Compare commits

...
32 Commits
Author SHA1 Message Date
hh 0399145791 style(frontend): soften the document glyph nudge to a hairline 2026-08-13 04:13:03 +02:00
hh 0b8e90159f style(frontend): optically nudge the document glyph in file chips 2026-08-13 04:07:34 +02:00
hh 995322b17a fix(frontend): readable links and file chips on own bubbles 2026-08-13 04:06:44 +02:00
hh 6ae28cb321 fix(frontend): render document albums as file rows instead of broken images 2026-08-13 03:59:01 +02:00
hh e833de245a fix(api,frontend): render files by name without fetching, previewable share links 2026-08-13 03:49:29 +02:00
hh ef739838a5 fix(migrations): lift timescale decompression limit when dropping scheduled dupes 2026-08-13 03:30:07 +02:00
hh 3e08698b62 feat(userbot,api,frontend): file links, media downloads, story hold-pause, drop scheduled dupes 2026-08-13 03:17:36 +02:00
hh 683b9a31a3 fix(frontend): restore outside-click dismiss, collapse story groups by default 2026-08-06 15:28:10 +02:00
hh 004fd56f2c feat(frontend): open sidebar media on click, unbreak long-press menus 2026-08-06 15:17:25 +02:00
hh 12cc7d57e3 feat(userbot,api,frontend): backfill stories and unbreak all-stories list 2026-08-06 14:08:10 +02:00
hh 1e143c7573 fix(frontend): unbreak round video ring and playback in firefox 2026-08-06 13:52:27 +02:00
hh ee63f8b783 feat(api,userbot,frontend): fetch only new messages in backfill by default 2026-08-06 12:30:37 +02:00
hh 6b6edc9a0d fix(frontend): paint video poster frames, clip round videos to circle 2026-08-06 12:29:37 +02:00
hh 525ce024bc perf(api,frontend): index hot queries, paginate chats, fix realtime 2026-08-06 01:52:41 +02:00
hh 9d767d2531 chore(frontend): sync bun lockfile 2026-08-06 01:07:52 +02:00
hh 18220407af fix(make): enable db and migrate profiles in migrate target 2026-08-06 01:07:51 +02:00
hh 1898a51a9d feat(api,userbot,frontend): rename session device from the web ui 2026-08-06 00:51:58 +02:00
hh b1848a6620 feat(caddy): move routes into site.caddy 2026-08-06 00:50:29 +02:00
hh e887dfc5ce feat(api,frontend): add qr code login 2026-08-06 00:23:38 +02:00
hh 9c265af3d3 feat(api,userbot,frontend): add accounts from the web ui and isolate per-account settings 2026-08-05 23:49:52 +02:00
hh 92fd20137e feat(api,userbot,frontend): search peers without chats and start tracking them 2026-08-05 23:46:54 +02:00
hh dcf95bd9d4 fix: adapt for mobile 2026-06-13 14:40:24 +02:00
hh 6ed392617f feat: add restart policy 2026-06-07 00:47:00 +02:00
hh 4fdf70a898 fix: mcp with token should work 2026-06-02 02:20:25 +02:00
hh f3712cfe36 fix: circular import and internal storage 2026-06-02 02:04:59 +02:00
hh f688530eac feat: proper frontend passing 2026-06-02 01:39:43 +02:00
hh 3aaa3c757f feat: add caddy config 2026-06-02 01:22:58 +02:00
hh 17cd31c41e feat: mcp and backfill fixes 2026-06-02 01:02:08 +02:00
hh c6984a7286 feat: add instructions to mcp 2026-06-01 23:45:54 +02:00
hh 2465bcd184 feat: add annotations, user profiles, watchers, stories, search and more 2026-06-01 17:15:09 +02:00
hh ed469ba8dd feat: web UI chat render, panels, presence + analytics 2026-05-31 19:41:37 +02:00
hh 75425d1bee feat: 1-to-1 message render + web data-lake backend 2026-05-31 01:45:05 +02:00
253 changed files with 28080 additions and 302 deletions
+3
View File
@@ -1,6 +1,9 @@
COMPOSE_PROFILES=db,userbot,api
RUN_ENVIRONMENT=prod
# Leave empty for *.localhost.
FRONTEND_DEV_HOST=
DB__HOST=postgres
DB__PORT=5432
DB__USER=beavergram
+6 -2
View File
@@ -1,4 +1,4 @@
.PHONY: recreate down restart rebuild deploy migrate session-create
.PHONY: recreate down restart rebuild deploy migrate session-create frontend
recreate:
docker compose up -d --force-recreate
@@ -13,11 +13,15 @@ rebuild:
docker compose build
docker compose up -d
frontend:
docker compose run --rm --no-deps frontend-dev sh -c "bun install && bun run build"
deploy:
$(MAKE) frontend
$(MAKE) rebuild
migrate:
docker compose run --rm migrator $(filter-out $@,$(MAKECMDGOALS))
docker compose --profile db --profile migrate run --rm migrator $(filter-out $@,$(MAKECMDGOALS))
session-create:
cd backend && uv run python scripts/session/create.py
@@ -0,0 +1,27 @@
"""media unique_id for edit detection
Revision ID: a9c3e7f1d2b4
Revises: f7a2c9e1b3d5
Create Date: 2026-05-30 04:00:00.000000
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "a9c3e7f1d2b4"
down_revision: str | None = "f7a2c9e1b3d5"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.add_column("media", sa.Column("unique_id", sa.String(), nullable=True))
op.add_column("media_versions", sa.Column("unique_id", sa.String(), nullable=True))
def downgrade() -> None:
op.drop_column("media_versions", "unique_id")
op.drop_column("media", "unique_id")
@@ -0,0 +1,24 @@
"""account device model
Revision ID: b9e4d1a70c26
Revises: e7b4c2a9f861
Create Date: 2026-08-06 01:20:00.000000
"""
from collections.abc import Sequence
from alembic import op
revision: str = "b9e4d1a70c26"
down_revision: str | None = "e7b4c2a9f861"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.execute("ALTER TABLE accounts ADD COLUMN device_model text")
def downgrade() -> None:
op.execute("ALTER TABLE accounts DROP COLUMN device_model")
@@ -0,0 +1,41 @@
"""hot path indexes
Revision ID: c1f6b3d84a92
Revises: b9e4d1a70c26
Create Date: 2026-08-06 12:00:00.000000
"""
from collections.abc import Sequence
from alembic import op
revision: str = "c1f6b3d84a92"
down_revision: str | None = "b9e4d1a70c26"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.execute(
"CREATE INDEX ix_messages_chat_date ON messages "
"(account_id, chat_id, date DESC, message_id DESC)"
)
op.execute("CREATE INDEX ix_avatars_owner ON avatars (account_id, owner_id)")
op.execute("CREATE INDEX ix_media_message ON media (account_id, message_id)")
op.execute(
"CREATE INDEX ix_chat_history_chat_ts ON chat_history "
"(account_id, chat_id, ts DESC)"
)
op.execute(
"CREATE INDEX ix_read_receipts_chat ON read_receipts "
"(account_id, chat_id, kind, message_id DESC)"
)
def downgrade() -> None:
op.execute("DROP INDEX IF EXISTS ix_read_receipts_chat")
op.execute("DROP INDEX IF EXISTS ix_chat_history_chat_ts")
op.execute("DROP INDEX IF EXISTS ix_media_message")
op.execute("DROP INDEX IF EXISTS ix_avatars_owner")
op.execute("DROP INDEX IF EXISTS ix_messages_chat_date")
@@ -0,0 +1,42 @@
"""custom emoji documents
Revision ID: c4e8a1f7d9b2
Revises: a9c3e7f1d2b4
Create Date: 2026-05-31 12:00:00.000000
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
revision: str = "c4e8a1f7d9b2"
down_revision: str | None = "a9c3e7f1d2b4"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"custom_emoji",
sa.Column("custom_emoji_id", sa.BigInteger(), nullable=False),
sa.Column("storage_key", sa.String(), nullable=True),
sa.Column("file_size", sa.BigInteger(), nullable=True),
sa.Column("mime", sa.String(), nullable=True),
sa.Column("kind", sa.String(), nullable=True),
sa.Column("downloaded", sa.Boolean(), nullable=False),
sa.Column("raw", postgresql.JSONB(astext_type=sa.Text()), nullable=False),
sa.Column(
"first_seen_at",
sa.DateTime(timezone=True),
server_default=sa.text("now()"),
nullable=False,
),
sa.PrimaryKeyConstraint("custom_emoji_id"),
)
def downgrade() -> None:
op.drop_table("custom_emoji")
@@ -0,0 +1,119 @@
"""chat stats
Revision ID: d4a7e2b91f38
Revises: c1f6b3d84a92
Create Date: 2026-08-06 12:30:00.000000
"""
from collections.abc import Sequence
from alembic import op
revision: str = "d4a7e2b91f38"
down_revision: str | None = "c1f6b3d84a92"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_APPLY = """
CREATE FUNCTION chat_stats_apply() RETURNS trigger LANGUAGE plpgsql AS $$
BEGIN
IF TG_OP = 'INSERT' THEN
INSERT INTO chat_stats AS cs (account_id, chat_id, message_count,
last_date, last_message_id,
last_text, last_sender_id)
VALUES (NEW.account_id, NEW.chat_id, 1,
CASE WHEN NEW.date <= now() + interval '1 day'
THEN NEW.date END,
CASE WHEN NEW.date <= now() + interval '1 day'
THEN NEW.message_id END,
NEW.text, NEW.sender_id)
ON CONFLICT (account_id, chat_id) DO UPDATE SET
message_count = cs.message_count + 1,
last_date = CASE WHEN chat_stats_newer(cs, EXCLUDED)
THEN EXCLUDED.last_date ELSE cs.last_date END,
last_message_id = CASE WHEN chat_stats_newer(cs, EXCLUDED)
THEN EXCLUDED.last_message_id
ELSE cs.last_message_id END,
last_text = CASE WHEN chat_stats_newer(cs, EXCLUDED)
THEN EXCLUDED.last_text ELSE cs.last_text END,
last_sender_id = CASE WHEN chat_stats_newer(cs, EXCLUDED)
THEN EXCLUDED.last_sender_id
ELSE cs.last_sender_id END;
ELSE
UPDATE chat_stats
SET last_text = NEW.text, last_sender_id = NEW.sender_id
WHERE account_id = NEW.account_id
AND chat_id = NEW.chat_id
AND last_message_id = NEW.message_id;
END IF;
RETURN NULL;
END;
$$
"""
_NEWER = """
CREATE FUNCTION chat_stats_newer(current chat_stats, incoming chat_stats)
RETURNS boolean LANGUAGE sql IMMUTABLE AS $$
SELECT incoming.last_date IS NOT NULL
AND (current.last_date IS NULL
OR (incoming.last_date, incoming.last_message_id)
> (current.last_date, current.last_message_id))
$$
"""
_BACKFILL = """
INSERT INTO chat_stats (account_id, chat_id, message_count, last_date,
last_message_id, last_text, last_sender_id)
SELECT agg.account_id, agg.chat_id, agg.message_count,
last.date, last.message_id, last.text, last.sender_id
FROM (
SELECT account_id, chat_id, count(*) AS message_count
FROM messages GROUP BY account_id, chat_id
) agg
LEFT JOIN LATERAL (
SELECT date, message_id, text, sender_id FROM messages m
WHERE m.account_id = agg.account_id AND m.chat_id = agg.chat_id
AND m.date <= now() + interval '1 day'
ORDER BY m.date DESC, m.message_id DESC LIMIT 1
) last ON true
"""
def upgrade() -> None:
op.execute(
"CREATE TABLE chat_stats ("
"account_id integer NOT NULL, "
"chat_id bigint NOT NULL, "
"message_count bigint NOT NULL DEFAULT 0, "
"last_date timestamptz, "
"last_message_id bigint, "
"last_text text, "
"last_sender_id bigint, "
"PRIMARY KEY (account_id, chat_id))"
)
op.execute(_BACKFILL)
op.execute(
"CREATE INDEX ix_chat_stats_recent ON chat_stats "
"(account_id, last_date DESC, chat_id DESC)"
)
op.execute(_NEWER)
op.execute(_APPLY)
op.execute(
"CREATE TRIGGER messages_chat_stats_insert AFTER INSERT ON messages "
"FOR EACH ROW EXECUTE FUNCTION chat_stats_apply()"
)
op.execute(
"CREATE TRIGGER messages_chat_stats_update AFTER UPDATE ON messages "
"FOR EACH ROW WHEN (OLD.text IS DISTINCT FROM NEW.text "
"OR OLD.sender_id IS DISTINCT FROM NEW.sender_id) "
"EXECUTE FUNCTION chat_stats_apply()"
)
def downgrade() -> None:
op.execute("DROP TRIGGER IF EXISTS messages_chat_stats_update ON messages")
op.execute("DROP TRIGGER IF EXISTS messages_chat_stats_insert ON messages")
op.execute("DROP FUNCTION IF EXISTS chat_stats_apply()")
op.execute("DROP FUNCTION IF EXISTS chat_stats_newer(chat_stats, chat_stats)")
op.execute("DROP TABLE IF EXISTS chat_stats")
@@ -0,0 +1,36 @@
"""dialogs
Revision ID: d5f9b2c8e3a1
Revises: c4e8a1f7d9b2
Create Date: 2026-05-31 18:00:00.000000
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "d5f9b2c8e3a1"
down_revision: str | None = "c4e8a1f7d9b2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"dialogs",
sa.Column("account_id", sa.BigInteger(), nullable=False),
sa.Column("chat_id", sa.BigInteger(), nullable=False),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
server_default=sa.text("now()"),
nullable=False,
),
sa.PrimaryKeyConstraint("account_id", "chat_id"),
)
def downgrade() -> None:
op.drop_table("dialogs")
@@ -0,0 +1,75 @@
"""drop scheduled message duplicates
Revision ID: e1c7a4b62d90
Revises: d4a7e2b91f38
Create Date: 2026-08-13 10:00:00.000000
"""
from collections.abc import Sequence
from alembic import op
revision: str = "e1c7a4b62d90"
down_revision: str | None = "d4a7e2b91f38"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_STAGE = """
CREATE TEMP TABLE scheduled_keys ON COMMIT DROP AS
SELECT DISTINCT s.account_id, s.chat_id, s.message_id
FROM messages s
WHERE s.raw->>'scheduled' = 'true'
AND NOT EXISTS (
SELECT 1 FROM messages r
WHERE r.account_id = s.account_id AND r.chat_id = s.chat_id
AND r.message_id = s.message_id
AND r.raw->>'scheduled' IS DISTINCT FROM 'true'
)
"""
_DELETE_CHILD = """
DELETE FROM {table} t
USING scheduled_keys k
WHERE t.account_id = k.account_id AND t.chat_id = k.chat_id
AND t.message_id = k.message_id
"""
_DELETE_MESSAGES = """
DELETE FROM messages m
USING scheduled_keys k
WHERE m.account_id = k.account_id AND m.chat_id = k.chat_id
AND m.message_id = k.message_id AND m.raw->>'scheduled' = 'true'
"""
_RECOUNT = """
UPDATE chat_stats cs
SET message_count = fresh.message_count
FROM (
SELECT k.account_id, k.chat_id,
(SELECT count(*) FROM messages m
WHERE m.account_id = k.account_id AND m.chat_id = k.chat_id)
AS message_count
FROM (SELECT DISTINCT account_id, chat_id FROM scheduled_keys) k
) fresh
WHERE cs.account_id = fresh.account_id AND cs.chat_id = fresh.chat_id
"""
_CHILD_TABLES = ("media", "media_versions", "message_versions", "links", "callbacks")
_ALLOW_DECOMPRESSION = (
"SET LOCAL timescaledb.max_tuples_decompressed_per_dml_transaction = 0"
)
def upgrade() -> None:
op.execute(_ALLOW_DECOMPRESSION)
op.execute(_STAGE)
for table in _CHILD_TABLES:
op.execute(_DELETE_CHILD.format(table=table))
op.execute(_DELETE_MESSAGES)
op.execute(_RECOUNT)
def downgrade() -> None:
pass
@@ -0,0 +1,43 @@
"""per-account policy defaults
Revision ID: e7b4c2a9f861
Revises: d5f9b2c8e3a1
Create Date: 2026-08-05 23:10:00.000000
"""
from collections.abc import Sequence
from alembic import op
revision: str = "e7b4c2a9f861"
down_revision: str | None = "d5f9b2c8e3a1"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_SCOPE_KEY = "capture_policy_account_id_scope_type_scope_id_key"
def upgrade() -> None:
op.execute("DROP INDEX ix_capture_policy_default")
op.execute(f"ALTER TABLE capture_policy DROP CONSTRAINT {_SCOPE_KEY}")
op.execute(
"CREATE UNIQUE INDEX ix_capture_policy_scope ON capture_policy "
"(account_id, scope_type, scope_id) NULLS NOT DISTINCT"
)
def downgrade() -> None:
op.execute(
"DELETE FROM capture_policy WHERE account_id IS NOT NULL "
"AND scope_type LIKE 'default_%'"
)
op.execute("DROP INDEX ix_capture_policy_scope")
op.execute(
f"ALTER TABLE capture_policy ADD CONSTRAINT {_SCOPE_KEY} "
"UNIQUE (account_id, scope_type, scope_id)"
)
op.execute(
"CREATE UNIQUE INDEX ix_capture_policy_default "
"ON capture_policy (scope_type) WHERE scope_type LIKE 'default_%'"
)
@@ -0,0 +1,73 @@
"""file shares
Revision ID: f2b8d3c9a51e
Revises: e1c7a4b62d90
Create Date: 2026-08-13 10:30:00.000000
"""
from collections.abc import Sequence
from alembic import op
revision: str = "f2b8d3c9a51e"
down_revision: str | None = "e1c7a4b62d90"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_SHARES = """
CREATE TABLE file_shares (
id serial PRIMARY KEY,
account_id integer NOT NULL,
token text NOT NULL UNIQUE,
kind text NOT NULL,
storage_key text NOT NULL,
file_name text NOT NULL,
mime text,
file_size bigint,
title text,
chat_id bigint,
message_id bigint,
peer_id bigint,
story_id bigint,
expires_at timestamptz,
max_downloads integer,
download_count integer NOT NULL DEFAULT 0,
last_download_at timestamptz,
revoked_at timestamptz,
created_at timestamptz NOT NULL DEFAULT now()
)
"""
_HITS = """
CREATE TABLE file_share_hits (
id bigserial PRIMARY KEY,
share_id integer NOT NULL REFERENCES file_shares (id) ON DELETE CASCADE,
ts timestamptz NOT NULL DEFAULT now(),
method text NOT NULL,
ip text,
user_agent text,
counted boolean NOT NULL DEFAULT false
)
"""
def upgrade() -> None:
op.execute(_SHARES)
op.execute(
"CREATE INDEX ix_file_shares_account ON file_shares "
"(account_id, created_at DESC)"
)
op.execute(
"CREATE INDEX ix_file_shares_subject ON file_shares "
"(account_id, storage_key) WHERE revoked_at IS NULL"
)
op.execute(_HITS)
op.execute(
"CREATE INDEX ix_file_share_hits_share ON file_share_hits (share_id, ts DESC)"
)
def downgrade() -> None:
op.execute("DROP TABLE IF EXISTS file_share_hits")
op.execute("DROP TABLE IF EXISTS file_shares")
@@ -0,0 +1,56 @@
"""media versions (keep every edited media)
Revision ID: f7a2c9e1b3d5
Revises: a3f1c8e94d72
Create Date: 2026-05-30 03:00:00.000000
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "f7a2c9e1b3d5"
down_revision: str | None = "a3f1c8e94d72"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"media_versions",
sa.Column("id", sa.BigInteger(), autoincrement=True, nullable=False),
sa.Column("account_id", sa.Integer(), nullable=False),
sa.Column("chat_id", sa.BigInteger(), nullable=False),
sa.Column("message_id", sa.BigInteger(), nullable=False),
sa.Column(
"observed_at",
sa.DateTime(timezone=True),
server_default=sa.text("now()"),
nullable=False,
),
sa.Column("kind", sa.String(), nullable=False),
sa.Column("storage_key", sa.String(), nullable=False),
sa.Column("file_size", sa.BigInteger(), nullable=True),
sa.Column("mime", sa.String(), nullable=True),
sa.Column("ttl_seconds", sa.Integer(), nullable=True),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint(
"account_id",
"chat_id",
"message_id",
"storage_key",
name="uq_media_versions_content",
),
)
op.create_index(
"ix_media_versions_message",
"media_versions",
["account_id", "chat_id", "message_id", "observed_at"],
)
def downgrade() -> None:
op.drop_index("ix_media_versions_message", table_name="media_versions")
op.drop_table("media_versions")
+50 -1
View File
@@ -1,28 +1,42 @@
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from pathlib import Path
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka, setup_dishka
from fastapi import FastAPI
from fastapi import FastAPI, Request
from fastapi.responses import FileResponse, RedirectResponse
from fastmcp.utilities.lifespan import combine_lifespans
from starlette.applications import Starlette
from api.auth import BearerAuthMiddleware
from api.mcp.server import mcp
from api.realtime import hub
from api.routers import (
accounts,
analytics,
annotations,
avatars,
backfill,
chats,
custom_emoji,
discover,
events,
files,
folders,
media,
peers,
policy,
presence,
profile,
search,
shares,
social,
stories,
watches,
)
from dependencies.container import container
from utils.cache import DAY_HEADERS, IMMUTABLE_HEADERS, NO_STORE_HEADERS
from utils.env import env
if env.auth.token is None:
@@ -35,7 +49,10 @@ mcp_app = mcp.http_app(path="/")
@asynccontextmanager
async def lifespan(app_: Starlette) -> AsyncGenerator[None]:
pool = await container.get(asyncpg.Pool)
await hub.start(pool)
yield
await hub.stop()
await app_.state.dishka_container.close()
@@ -54,20 +71,52 @@ async def health(pool: FromDishka[asyncpg.Pool]) -> dict[str, bool]:
return {"db": db_ok, "timescaledb": bool(timescale_ok)}
app.include_router(accounts.router)
app.include_router(analytics.router)
app.include_router(policy.router)
app.include_router(folders.router)
app.include_router(backfill.router)
app.include_router(search.router)
app.include_router(chats.router)
app.include_router(media.router)
app.include_router(avatars.router)
app.include_router(custom_emoji.router)
app.include_router(social.router)
app.include_router(presence.router)
app.include_router(stories.router)
app.include_router(profile.router)
app.include_router(events.router)
app.include_router(peers.router)
app.include_router(discover.router)
app.include_router(annotations.router)
app.include_router(watches.router)
app.include_router(shares.router)
app.include_router(files.router)
app.mount("/mcp", mcp_app)
@app.api_route("/mcp", methods=["GET", "POST", "DELETE"])
async def mcp_trailing_slash(request: Request) -> RedirectResponse:
query = request.url.query
return RedirectResponse(f"/mcp/?{query}" if query else "/mcp/", status_code=307)
_spa_dir = Path(env.api.static_dir).resolve()
if _spa_dir.is_dir():
_spa_index = _spa_dir / "index.html"
@app.get("/{spa_path:path}")
async def serve_spa(spa_path: str) -> FileResponse:
candidate = (_spa_dir / spa_path).resolve()
if spa_path and candidate.is_relative_to(_spa_dir) and candidate.is_file():
immutable = spa_path.startswith("_app/immutable/")
return FileResponse(
candidate, headers=IMMUTABLE_HEADERS if immutable else DAY_HEADERS
)
return FileResponse(_spa_index, headers=NO_STORE_HEADERS)
app.add_middleware(BearerAuthMiddleware, token=_token)
setup_dishka(container, app)
+13 -3
View File
@@ -1,3 +1,6 @@
from secrets import compare_digest
from urllib.parse import parse_qs
from starlette.types import ASGIApp, Receive, Scope, Send
PROTECTED_PREFIXES = ("/api", "/mcp")
@@ -9,6 +12,15 @@ class BearerAuthMiddleware:
self.app = app
self.token = token
def _authorized(self, scope: Scope) -> bool:
headers = dict(scope["headers"])
bearer = headers.get(b"authorization", b"").decode()
if bearer.startswith("Bearer ") and compare_digest(bearer[7:], self.token):
return True
query = parse_qs(scope["query_string"].decode())
token = query.get("token", [""])[0]
return bool(token) and compare_digest(token, self.token)
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
@@ -18,9 +30,7 @@ class BearerAuthMiddleware:
):
await self.app(scope, receive, send)
return
headers = dict(scope["headers"])
authorization = headers.get(b"authorization", b"").decode()
if authorization == f"Bearer {self.token}":
if self._authorized(scope):
await self.app(scope, receive, send)
return
await send(
+183
View File
@@ -0,0 +1,183 @@
import asyncio
import contextlib
import secrets
import time
from dataclasses import dataclass, field
from pathlib import Path
from pyrogram.errors import AuthTokenExpired, SessionPasswordNeeded
from pyrogram.qrlogin import QRLogin
from pyrogram.types import User
from userbot import PyroClient
from utils.env import env
LOGIN_TTL_SECONDS = 900
QR_POLL_SECONDS = 25
PENDING_DIRNAME = "pending"
class LoginError(Exception):
pass
@dataclass
class QrState:
url: str
user: User | None = None
password_needed: bool = False
error: str | None = None
changed: asyncio.Event = field(default_factory=asyncio.Event)
@property
def settled(self) -> bool:
return self.user is not None or self.password_needed or self.error is not None
@dataclass
class PendingLogin:
client: PyroClient
started_at: float
phone: str = ""
phone_code_hash: str = ""
qr: QrState | None = None
watcher: asyncio.Task[None] | None = None
def _sessions_dir() -> Path:
path = Path(env.tg.sessions_dir)
path.mkdir(parents=True, exist_ok=True)
return path
def _pending_dir() -> Path:
path = _sessions_dir() / PENDING_DIRNAME
path.mkdir(parents=True, exist_ok=True)
return path
async def _stop_watcher(login: PendingLogin) -> None:
if login.watcher is None or login.watcher.done():
return
login.watcher.cancel()
with contextlib.suppress(asyncio.CancelledError, Exception):
await login.watcher
async def _watch_qr(qr: QRLogin, state: QrState) -> None:
while not state.settled:
try:
try:
state.user = await qr.wait()
except (TimeoutError, AuthTokenExpired):
await qr.recreate()
state.url = qr.url
except SessionPasswordNeeded:
state.password_needed = True
except Exception as exc:
state.error = str(exc)
state.changed.set()
class LoginManager:
def __init__(self) -> None:
self._logins: dict[str, PendingLogin] = {}
def _get(self, login_id: str) -> PendingLogin:
login = self._logins.get(login_id)
if login is None:
msg = "Сессия входа истекла, начните заново"
raise LoginError(msg)
return login
async def _connect(self) -> tuple[str, PyroClient]:
await self._sweep()
login_id = secrets.token_hex(8)
client = PyroClient(login_id, workdir=str(_pending_dir()), load_handlers=False)
await client.connect()
return login_id, client
async def start(self, phone: str) -> str:
login_id, client = await self._connect()
try:
sent = await client.send_code(phone)
except Exception:
await self._discard(login_id, client)
raise
self._logins[login_id] = PendingLogin(
client, time.monotonic(), phone=phone, phone_code_hash=sent.phone_code_hash
)
return login_id
async def start_qr(self) -> tuple[str, str]:
login_id, client = await self._connect()
qr = QRLogin(client)
try:
await qr.recreate()
except Exception:
await self._discard(login_id, client)
raise
state = QrState(qr.url)
self._logins[login_id] = PendingLogin(
client,
time.monotonic(),
qr=state,
watcher=asyncio.create_task(_watch_qr(qr, state)),
)
return login_id, state.url
async def wait_qr(self, login_id: str) -> QrState:
login = self._get(login_id)
if login.qr is None:
msg = "Этот вход начат по номеру телефона"
raise LoginError(msg)
with contextlib.suppress(TimeoutError):
async with asyncio.timeout(QR_POLL_SECONDS):
await login.qr.changed.wait()
login.qr.changed.clear()
return login.qr
async def submit_code(self, login_id: str, code: str) -> User | None:
login = self._get(login_id)
try:
user = await login.client.sign_in(login.phone, login.phone_code_hash, code)
except SessionPasswordNeeded:
return None
if not isinstance(user, User):
await self.cancel(login_id)
msg = "Этот номер не зарегистрирован в Telegram"
raise LoginError(msg)
return user
async def submit_password(self, login_id: str, password: str) -> User:
return await self._get(login_id).client.check_password(password)
async def finalize(self, login_id: str, session_name: str) -> None:
login = self._logins.pop(login_id)
await _stop_watcher(login)
await login.client.disconnect()
source = _pending_dir() / f"{login_id}.session"
source.replace(_sessions_dir() / f"{session_name}.session")
async def cancel(self, login_id: str) -> None:
login = self._logins.pop(login_id, None)
if login is not None:
await _stop_watcher(login)
await self._discard(login_id, login.client)
async def _discard(self, login_id: str, client: PyroClient) -> None:
with contextlib.suppress(Exception):
await client.disconnect()
(_pending_dir() / f"{login_id}.session").unlink(missing_ok=True)
async def _sweep(self) -> None:
now = time.monotonic()
for login_id, login in list(self._logins.items()):
if now - login.started_at > LOGIN_TTL_SECONDS:
await self.cancel(login_id)
for path in _pending_dir().glob("*.session"):
if path.stem not in self._logins:
path.unlink(missing_ok=True)
login_manager = LoginManager()
+173
View File
@@ -0,0 +1,173 @@
from datetime import datetime
from typing import Any
import asyncpg
from fastmcp.utilities.types import Image
from mcp.types import TextContent
from utils.read import peers
from utils.read.models import MediaRef, MessageView
from utils.storage import ContentAddressedStorage
_VOICE_KINDS = {"voice", "video_note"}
def _ts(value: datetime) -> str:
return value.strftime("%Y-%m-%d %H:%M")
def _name(sender_id: int | None, names: dict[int, str], self_id: int | None) -> str:
if sender_id is None:
return "Unknown"
if sender_id == self_id:
return "Me"
return names.get(sender_id) or str(sender_id)
def _media_note(media: list[MediaRef]) -> list[str]:
notes: list[str] = []
for item in media:
if item.kind in _VOICE_KINDS:
if item.extracted_text:
notes.append(f"(Voice message, STT Content: {item.extracted_text})")
else:
notes.append("(Voice message, not transcribed)")
elif item.kind == "photo":
state = "" if item.downloaded else ", not downloaded"
notes.append(f"[photo #{item.message_id}{state}]")
else:
notes.append(f"[{item.kind}]")
return notes
def _line(
view: MessageView,
names: dict[int, str],
self_id: int | None,
notes: dict[int, list[str]],
) -> str:
parts: list[str] = []
if view.reply and (view.reply.sender_name or view.reply.text):
ref = view.reply.sender_name or str(view.reply.sender_id or "?")
parts.append(f"(reply to {ref})")
if view.text:
parts.append(view.text)
parts.extend(_media_note(view.media))
suffix = ""
if view.edited_at:
suffix += " (edited)"
if view.deleted_at:
suffix += " (deleted)"
body = " ".join(part for part in parts if part) or "(no text)"
name = _name(view.sender_id, names, self_id)
line = f"#{view.message_id} {name} ({_ts(view.date)}): {body}{suffix}"
for note in notes.get(view.message_id, []):
line += f"\n 📝 [your private note, NOT in Telegram]: {note}"
return line
async def load_notes(
pool: asyncpg.Pool, account_id: int, chat_id: int, views: list[MessageView]
) -> dict[int, list[str]]:
ids = [view.message_id for view in views]
if not ids:
return {}
rows = await pool.fetch(
"SELECT message_id, text FROM annotations "
"WHERE account_id = $1 AND chat_id = $2 AND message_id = ANY($3::bigint[]) "
"ORDER BY created_at",
account_id,
chat_id,
ids,
)
notes: dict[int, list[str]] = {}
for row in rows:
notes.setdefault(row["message_id"], []).append(row["text"])
return notes
async def resolve_names(
pool: asyncpg.Pool, account_id: int, views: list[MessageView]
) -> dict[int, str]:
ids = list({view.sender_id for view in views if view.sender_id is not None})
found = await peers.get_peers(pool, account_id, ids)
names: dict[int, str] = {}
for peer in found:
name = (
" ".join(part for part in (peer.first_name, peer.last_name) if part)
or peer.username
)
if name:
names[peer.peer_id] = name
return names
async def load_photos(
pool: asyncpg.Pool,
storage: ContentAddressedStorage,
account_id: int,
views: list[MessageView],
*,
limit: int,
) -> tuple[list[tuple[int, bytes, str]], bool]:
refs = [
(item.message_id, item.id)
for view in views
for item in view.media
if item.kind == "photo" and item.downloaded and item.id is not None
]
truncated = len(refs) > limit
refs = refs[-limit:]
if not refs:
return [], truncated
rows = await pool.fetch(
"SELECT id, storage_key, mime FROM media "
"WHERE account_id = $1 AND id = ANY($2::bigint[])",
account_id,
[media_id for _, media_id in refs],
)
by_id = {row["id"]: row for row in rows}
out: list[tuple[int, bytes, str]] = []
for message_id, media_id in refs:
row = by_id.get(media_id)
if row is None or not row["storage_key"]:
continue
try:
data = storage.get(row["storage_key"])
except OSError:
continue
fmt = (row["mime"] or "image/jpeg").split("/")[-1]
out.append((message_id, data, fmt))
return out, truncated
def build_transcript(
views: list[MessageView],
names: dict[int, str],
self_id: int | None,
photos: list[tuple[int, bytes, str]],
*,
notes: dict[int, list[str]] | None = None,
truncated: bool = False,
) -> list[Any]:
if not views:
return [TextContent(type="text", text="No messages.")]
notes = notes or {}
header = f"{len(views)} messages (oldest first)"
body = f"{header}\n\n" + "\n".join(
_line(view, names, self_id, notes) for view in views
)
blocks: list[Any] = [TextContent(type="text", text=body)]
if truncated:
blocks.append(
TextContent(
type="text",
text="(images truncated to the most recent; narrow the range for more)",
)
)
for message_id, data, fmt in photos:
blocks.append(
TextContent(type="text", text=f"Image attached to message #{message_id}:")
)
blocks.append(Image(data=data, format=fmt))
return blocks
+111 -12
View File
@@ -6,14 +6,57 @@ import asyncpg
from fastmcp import FastMCP
from pydantic import BaseModel
from api.mcp.format import build_transcript, load_notes, load_photos, resolve_names
from dependencies.container import container
from utils.jobs import enqueue
from utils.read import annotations, chats, media, peers, presence, social, watches
from utils.read.accounts import self_user_id
from utils.read.models import DEFAULT_LIMIT, Page
from utils.search.models import SearchFilters
from utils.search.repository import search_messages
from utils.storage import ContentAddressedStorage
mcp: FastMCP = FastMCP("beavergram")
INSTRUCTIONS = """\
beavergram archives Telegram data (chats, messages, media, presence, stories,
peer history) into Postgres and exposes it read-only over these tools.
Account scoping:
- Every tool takes `account_id`. It selects which archived Telegram account to
read. Unless the user names a different one, always pass `account_id=1`.
Identifiers:
- `chat_id` and `peer_id` are Telegram IDs (negative for groups/channels,
positive for users/bots). Discover them with `list_chats` before calling
tools that need a specific chat.
Typical flows:
- Read a chat: `list_chats` -> pick a `chat_id` -> `get_chat_history` for a
human-readable transcript (names, timestamps, voice STT, inline photos), or
`get_chat_history_raw` for structured JSON. Both return messages oldest
first; default gives the latest `limit`. Page to older messages with
`before_id` (first id you saw), to newer with `after_id` (last id you saw);
begin a full forward walk with `after_id=0`. Each transcript line starts
with its `#message_id`, and photos show as `[photo #message_id]` — pass that
id to `get_media` (with `fetch=True` to download) or to annotation tools.
Lines marked
"📝 [your private note, NOT in Telegram]" are the user's own annotations
attached locally via the web UI; they never existed in Telegram. Treat them
as private notes from the user to you, not as chat content.
- Find something: `search_messages_tool` (full-text over message text and STT
transcripts; supports chat/sender/date filters and regex).
- Forensics: `get_deleted_messages` and `get_message_versions` recover content
removed or edited in Telegram but kept in the archive.
- Media: `get_media` returns metadata; pass `fetch=True` to enqueue a lazy
download if the file isn't stored yet.
- Monitoring: `set_watch` creates a local rule, `list_watches` lists rules, and
`list_alerts` reads what those rules fired.
Writes: everything is read-only except `set_watch`, the only allowed write.
Paging: list/search tools accept `limit` and `offset`.
"""
mcp: FastMCP = FastMCP("beavergram", instructions=INSTRUCTIONS)
async def _pool() -> asyncpg.Pool:
@@ -62,24 +105,80 @@ async def list_chats(
return _dump(await chats.list_chats(await _pool(), account_id, page))
@mcp.tool
@mcp.tool(output_schema=None)
async def get_chat_history(
account_id: int,
chat_id: int,
limit: int = DEFAULT_LIMIT,
offset: int = 0,
before_id: int | None = None,
after_id: int | None = None,
include_deleted: bool = True,
include_images: bool = True,
max_images: int = 20,
) -> list[Any]:
"""Read a chat as a readable transcript, oldest first.
Renders "Name (time): text", voice notes as
"(Voice message, STT Content: ...)", and inlines downloaded photos as
images. Default returns the latest `limit` messages; page to older
messages with `before_id` (the first id you saw) or to newer ones with
`after_id` (the last id you saw).
"""
pool = await _pool()
views = await chats.get_chat_history(
pool,
account_id,
chat_id,
Page(limit=limit),
include_deleted=include_deleted,
before_id=before_id,
after_id=after_id,
)
if after_id is None:
views = list(reversed(views))
names = await resolve_names(pool, account_id, views)
notes = await load_notes(pool, account_id, chat_id, views)
self_id = await self_user_id(pool, account_id)
photos: list[tuple[int, bytes, str]] = []
truncated = False
if include_images:
storage = await container.get(ContentAddressedStorage)
photos, truncated = await load_photos(
pool, storage, account_id, views, limit=max_images
)
return build_transcript(
views, names, self_id, photos, notes=notes, truncated=truncated
)
@mcp.tool
async def get_chat_history_raw(
account_id: int,
chat_id: int,
limit: int = DEFAULT_LIMIT,
before_id: int | None = None,
after_id: int | None = None,
include_deleted: bool = True,
) -> list[dict[str, Any]]:
"""Read archived messages of a chat, newest first."""
return _dump(
await chats.get_chat_history(
await _pool(),
account_id,
chat_id,
Page(limit=limit, offset=offset),
include_deleted=include_deleted,
)
"""Structured chat messages as JSON, oldest first.
Default returns the latest `limit` messages. Walk the whole chat forward
with `after_id` set to the last returned message_id (begin at
`after_id=0`), or backward with `before_id` set to the first returned
message_id.
"""
views = await chats.get_chat_history(
await _pool(),
account_id,
chat_id,
Page(limit=limit),
include_deleted=include_deleted,
before_id=before_id,
after_id=after_id,
)
if after_id is None:
views = list(reversed(views))
return _dump(views)
@mcp.tool
+128
View File
@@ -0,0 +1,128 @@
import asyncio
import json
import logging
from typing import Any
import asyncpg
from utils.env import env
from utils.events import BG_EVENTS_CHANNEL
from utils.read import chats as chats_read
from utils.read import presence as presence_read
logger = logging.getLogger(__name__)
QUEUE_MAXSIZE = 256
class Subscriber:
def __init__(self, account_id: int, chat_id: int | None) -> None:
self.account_id = account_id
self.chat_id = chat_id
self.queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue(maxsize=QUEUE_MAXSIZE)
class EventHub:
def __init__(self) -> None:
self._subscribers: set[Subscriber] = set()
self._pool: asyncpg.Pool | None = None
self._conn: asyncpg.Connection | None = None
self._tasks: set[asyncio.Task] = set()
def subscribe(self, account_id: int, chat_id: int | None) -> Subscriber:
sub = Subscriber(account_id, chat_id)
self._subscribers.add(sub)
return sub
def unsubscribe(self, sub: Subscriber) -> None:
self._subscribers.discard(sub)
async def start(self, pool: asyncpg.Pool) -> None:
self._pool = pool
conn = await asyncpg.connect(dsn=env.db.connection_url)
await conn.add_listener(BG_EVENTS_CHANNEL, self._on_notify)
self._conn = conn
logger.info("Realtime hub listening on %s", BG_EVENTS_CHANNEL)
async def stop(self) -> None:
for task in self._tasks:
task.cancel()
if self._conn is not None:
await self._conn.close()
self._conn = None
def _on_notify(
self, _conn: asyncpg.Connection, _pid: int, _channel: str, payload: str
) -> None:
task = asyncio.create_task(self._dispatch(payload))
self._tasks.add(task)
task.add_done_callback(self._tasks.discard)
async def _dispatch(self, payload: str) -> None:
try:
event = json.loads(payload)
except json.JSONDecodeError:
return
account_id = event.get("account_id")
chat_id = event.get("chat_id")
scoped = event.get("kind") == "presence"
targets = [
sub
for sub in self._subscribers
if sub.account_id == account_id
and (
sub.chat_id == chat_id
if scoped
else sub.chat_id is None or sub.chat_id == chat_id
)
]
if not targets:
return
frame = await self._build_frame(event)
if frame is None:
return
for sub in targets:
try:
sub.queue.put_nowait(frame)
except asyncio.QueueFull:
logger.warning("Dropping event for slow subscriber")
async def _build_frame( # noqa: PLR0911
self, event: dict[str, Any]
) -> dict[str, Any] | None:
if self._pool is None:
return None
kind = event.get("kind")
account_id = event["account_id"]
if kind in {"message", "edit", "reaction"}:
view = await chats_read.get_message(
self._pool, account_id, event["chat_id"], event["message_id"]
)
if view is None:
return None
return {"type": kind, "message": view.model_dump(mode="json")}
if kind == "delete":
return {
"type": "delete",
"chat_id": event.get("chat_id"),
"message_ids": event.get("message_ids", []),
}
if kind == "presence":
sample = await presence_read.current_presence(
self._pool, account_id, event["chat_id"]
)
return {
"type": "presence",
"peer_id": event["chat_id"],
"sample": sample.model_dump(mode="json") if sample else None,
}
if kind == "receipt":
return {
"type": "receipt",
"chat_id": event["chat_id"],
"read_up_to": event["message_id"],
}
return None
hub = EventHub()
+150
View File
@@ -0,0 +1,150 @@
from collections.abc import Coroutine
from typing import Any, Literal
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, status
from pydantic import BaseModel
from pyrogram.errors import FloodWait, RPCError
from pyrogram.types import User
from api.login import LoginError, login_manager
from userbot import DEVICE_MODEL_LIMIT
from utils.read import accounts
from utils.read.models import AccountView
router = APIRouter(prefix="/api", tags=["accounts"], route_class=DishkaRoute)
_ERROR_MESSAGES = {
"PhoneNumberInvalid": "Неверный номер телефона",
"PhoneNumberBanned": "Номер заблокирован в Telegram",
"PhoneCodeInvalid": "Неверный код",
"PhoneCodeExpired": "Код истёк, запросите новый",
"PasswordHashInvalid": "Неверный пароль",
}
class PhoneRequest(BaseModel):
phone: str
class CodeRequest(BaseModel):
code: str
class PasswordRequest(BaseModel):
password: str
class DeviceModelRequest(BaseModel):
device_model: str
class LoginState(BaseModel):
login_id: str
stage: Literal["code", "qr", "password", "done"]
qr_url: str | None = None
account: AccountView | None = None
async def _guard[T](coro: Coroutine[Any, Any, T]) -> T:
try:
return await coro
except LoginError as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(exc)) from exc
except FloodWait as exc:
raise HTTPException(
status.HTTP_429_TOO_MANY_REQUESTS,
f"Слишком много попыток, подождите {exc.value} с",
) from exc
except RPCError as exc:
detail = _ERROR_MESSAGES.get(type(exc).__name__, str(exc))
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail) from exc
async def _complete(pool: asyncpg.Pool, login_id: str, user: User) -> LoginState:
session_name = str(user.id)
account_id = await accounts.sync_account(pool, user, session_name)
await login_manager.finalize(login_id, session_name)
await accounts.notify_accounts_changed(pool)
return LoginState(
login_id=login_id,
stage="done",
account=await accounts.get_account(pool, account_id),
)
@router.get("/accounts")
async def list_accounts(pool: FromDishka[asyncpg.Pool]) -> list[AccountView]:
return await accounts.list_accounts(pool)
@router.post("/accounts/login")
async def start_login(body: PhoneRequest) -> LoginState:
login_id = await _guard(login_manager.start(body.phone))
return LoginState(login_id=login_id, stage="code")
@router.post("/accounts/login/qr")
async def start_qr_login() -> LoginState:
login_id, url = await _guard(login_manager.start_qr())
return LoginState(login_id=login_id, stage="qr", qr_url=url)
@router.get("/accounts/login/{login_id}/qr")
async def poll_qr_login(login_id: str, pool: FromDishka[asyncpg.Pool]) -> LoginState:
state = await _guard(login_manager.wait_qr(login_id))
if state.user is not None:
return await _complete(pool, login_id, state.user)
if state.password_needed:
return LoginState(login_id=login_id, stage="password")
if state.error is not None:
raise HTTPException(status.HTTP_400_BAD_REQUEST, state.error)
return LoginState(login_id=login_id, stage="qr", qr_url=state.url)
@router.post("/accounts/login/{login_id}/code")
async def submit_code(
login_id: str, body: CodeRequest, pool: FromDishka[asyncpg.Pool]
) -> LoginState:
user = await _guard(login_manager.submit_code(login_id, body.code))
if user is None:
return LoginState(login_id=login_id, stage="password")
return await _complete(pool, login_id, user)
@router.post("/accounts/login/{login_id}/password")
async def submit_password(
login_id: str, body: PasswordRequest, pool: FromDishka[asyncpg.Pool]
) -> LoginState:
user = await _guard(login_manager.submit_password(login_id, body.password))
return await _complete(pool, login_id, user)
@router.delete("/accounts/login/{login_id}", status_code=status.HTTP_204_NO_CONTENT)
async def cancel_login(login_id: str) -> None:
await login_manager.cancel(login_id)
@router.patch("/accounts/{account_id}/device")
async def rename_device(
account_id: int, body: DeviceModelRequest, pool: FromDishka[asyncpg.Pool]
) -> AccountView:
device_model = " ".join(body.device_model.split())
if not device_model or len(device_model) > DEVICE_MODEL_LIMIT:
raise HTTPException(
status.HTTP_400_BAD_REQUEST,
f"Имя устройства: от 1 до {DEVICE_MODEL_LIMIT} символов",
)
account = await accounts.set_device_model(pool, account_id, device_model)
if account is None:
raise HTTPException(status.HTTP_404_NOT_FOUND, "Аккаунт не найден")
await accounts.notify_accounts_changed(pool)
return account
@router.delete("/accounts/{account_id}", status_code=status.HTTP_204_NO_CONTENT)
async def logout_account(account_id: int, pool: FromDishka[asyncpg.Pool]) -> None:
if await accounts.deactivate_account(pool, account_id) is None:
raise HTTPException(status.HTTP_404_NOT_FOUND, "Аккаунт не найден")
await accounts.notify_accounts_changed(pool)
+30
View File
@@ -0,0 +1,30 @@
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, Query
from utils.read import analytics
from utils.read.models import ResponseStats, VolumeBucket
router = APIRouter(prefix="/api/analytics", tags=["analytics"], route_class=DishkaRoute)
AccountId = Annotated[int, Query()]
ChatId = Annotated[int, Query()]
@router.get("/volume")
async def volume(
pool: FromDishka[asyncpg.Pool],
account_id: AccountId,
chat_id: ChatId,
days: Annotated[int, Query()] = 90,
) -> list[VolumeBucket]:
return await analytics.message_volume(pool, account_id, chat_id, days=days)
@router.get("/response-time")
async def response_time(
pool: FromDishka[asyncpg.Pool], account_id: AccountId, chat_id: ChatId
) -> ResponseStats:
return await analytics.response_stats(pool, account_id, chat_id)
+59
View File
@@ -0,0 +1,59 @@
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, Query
from fastapi.responses import FileResponse
from utils.cache import IMMUTABLE_HEADERS, SHORT_HEADERS
from utils.jobs import enqueue
from utils.read.avatars import avatar_by_unique_id, avatar_history, current_avatar
from utils.read.models import AvatarHistoryView
from utils.storage import ContentAddressedStorage
router = APIRouter(prefix="/api/avatars", tags=["avatars"], route_class=DishkaRoute)
@router.get("/{owner_kind}/{owner_id}/history")
async def serve_avatar_history(
pool: FromDishka[asyncpg.Pool],
owner_kind: str, # noqa: ARG001
owner_id: int,
account_id: Annotated[int, Query()],
) -> list[AvatarHistoryView]:
return await avatar_history(pool, account_id, owner_id)
@router.get("/{owner_kind}/{owner_id}")
async def serve_avatar(
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
owner_kind: str,
owner_id: int,
account_id: Annotated[int, Query()],
unique_id: Annotated[str | None, Query()] = None,
) -> FileResponse:
avatar = (
await avatar_by_unique_id(pool, account_id, owner_id, unique_id)
if unique_id is not None
else await current_avatar(pool, account_id, owner_kind, owner_id)
)
if avatar is None:
raise HTTPException(status_code=404, detail="avatar not found")
if not avatar.downloaded or avatar.storage_key is None:
await enqueue(
pool,
account_id,
"fetch_avatar",
{
"owner_kind": owner_kind,
"owner_id": owner_id,
"unique_id": avatar.unique_id,
},
)
raise HTTPException(status_code=409, detail="avatar not downloaded; fetching")
return FileResponse(
storage.url(avatar.storage_key),
media_type=avatar.mime or "image/jpeg",
headers=IMMUTABLE_HEADERS if unique_id is not None else SHORT_HEADERS,
)
+63 -1
View File
@@ -16,6 +16,12 @@ class BackfillRequest(BaseModel):
account_id: int
chat_id: int
media: bool = False
full: bool = False
class StoriesBackfillRequest(BaseModel):
account_id: int
peer_id: int
class FetchMediaRequest(BaseModel):
@@ -24,6 +30,16 @@ class FetchMediaRequest(BaseModel):
message_id: int
class TranscribeRequest(BaseModel):
account_id: int
chat_id: int
message_id: int
class SyncDialogsRequest(BaseModel):
account_id: int
class EnqueueResponse(BaseModel):
job_id: int
@@ -60,7 +76,17 @@ async def enqueue_backfill(
pool,
body.account_id,
"backfill",
{"chat_id": body.chat_id, "media": body.media},
{"chat_id": body.chat_id, "media": body.media, "full": body.full},
)
return EnqueueResponse(job_id=job_id)
@router.post("/stories/backfill", status_code=201)
async def enqueue_stories_backfill(
pool: FromDishka[asyncpg.Pool], body: StoriesBackfillRequest
) -> EnqueueResponse:
job_id = await enqueue(
pool, body.account_id, "backfill_stories", {"peer_id": body.peer_id}
)
return EnqueueResponse(job_id=job_id)
@@ -78,6 +104,27 @@ async def enqueue_fetch_media(
return EnqueueResponse(job_id=job_id)
@router.post("/media/transcribe", status_code=201)
async def enqueue_transcribe(
pool: FromDishka[asyncpg.Pool], body: TranscribeRequest
) -> EnqueueResponse:
job_id = await enqueue(
pool,
body.account_id,
"transcribe",
{"chat_id": body.chat_id, "message_id": body.message_id},
)
return EnqueueResponse(job_id=job_id)
@router.post("/dialogs/sync", status_code=201)
async def enqueue_sync_dialogs(
pool: FromDishka[asyncpg.Pool], body: SyncDialogsRequest
) -> EnqueueResponse:
job_id = await enqueue(pool, body.account_id, "sync_dialogs", {})
return EnqueueResponse(job_id=job_id)
@router.get("/jobs")
async def list_jobs(
pool: FromDishka[asyncpg.Pool],
@@ -103,3 +150,18 @@ async def get_job(pool: FromDishka[asyncpg.Pool], job_id: int) -> JobView:
if row is None:
raise HTTPException(status_code=404, detail="job not found")
return _to_view(row)
@router.post("/jobs/{job_id}/cancel")
async def cancel_job(pool: FromDishka[asyncpg.Pool], job_id: int) -> JobView:
row = await pool.fetchrow(
"UPDATE jobs SET status = 'canceled', finished_at = now(), "
"updated_at = now() WHERE id = $1 AND status IN ('pending', 'running') "
"RETURNING *",
job_id,
)
if row is not None:
return _to_view(row)
if await pool.fetchval("SELECT 1 FROM jobs WHERE id = $1", job_id) is None:
raise HTTPException(status_code=404, detail="job not found")
raise HTTPException(status_code=409, detail="job already finished")
+46 -1
View File
@@ -3,7 +3,10 @@ from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, Query
from pydantic import BaseModel
from utils.jobs import enqueue
from utils.policy import repository
from utils.read import chats
from utils.read.models import (
DEFAULT_LIMIT,
@@ -11,10 +14,17 @@ from utils.read.models import (
MessageVersionView,
MessageView,
Page,
PinnedView,
)
from utils.read.pinned import get_pinned
router = APIRouter(prefix="/api", tags=["chats"], route_class=DishkaRoute)
class EnrichRequest(BaseModel):
account_id: int
AccountId = Annotated[int, Query()]
Limit = Annotated[int, Query()]
Offset = Annotated[int, Query()]
@@ -26,8 +36,24 @@ async def list_chats(
account_id: AccountId,
limit: Limit = DEFAULT_LIMIT,
offset: Offset = 0,
folder_id: Annotated[int | None, Query()] = None,
search: Annotated[str | None, Query()] = None,
) -> list[ChatListItem]:
return await chats.list_chats(pool, account_id, Page(limit=limit, offset=offset))
folder = (
await repository.get_folder(pool, account_id, folder_id)
if folder_id is not None
else None
)
return await chats.list_chats(
pool, account_id, Page(limit=limit, offset=offset), folder=folder, search=search
)
@router.get("/chats/{chat_id}")
async def get_chat(
pool: FromDishka[asyncpg.Pool], chat_id: int, account_id: AccountId
) -> ChatListItem | None:
return await chats.get_chat(pool, account_id, chat_id)
@router.get("/chats/{chat_id}/messages")
@@ -38,6 +64,8 @@ async def chat_history(
limit: Limit = DEFAULT_LIMIT,
offset: Offset = 0,
include_deleted: Annotated[bool, Query()] = True,
before_id: Annotated[int | None, Query()] = None,
after_id: Annotated[int | None, Query()] = None,
) -> list[MessageView]:
return await chats.get_chat_history(
pool,
@@ -45,9 +73,26 @@ async def chat_history(
chat_id,
Page(limit=limit, offset=offset),
include_deleted=include_deleted,
before_id=before_id,
after_id=after_id,
)
@router.get("/chats/{chat_id}/pinned")
async def chat_pinned(
pool: FromDishka[asyncpg.Pool], chat_id: int, account_id: AccountId
) -> PinnedView | None:
return await get_pinned(pool, account_id, chat_id)
@router.post("/chats/{chat_id}/enrich")
async def enrich_chat(
pool: FromDishka[asyncpg.Pool], chat_id: int, body: EnrichRequest
) -> dict[str, int]:
job_id = await enqueue(pool, body.account_id, "enrich_chat", {"chat_id": chat_id})
return {"job_id": job_id}
@router.get("/chats/{chat_id}/messages/{message_id}/versions")
async def message_versions(
pool: FromDishka[asyncpg.Pool], chat_id: int, message_id: int, account_id: AccountId
+35
View File
@@ -0,0 +1,35 @@
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, Query
from fastapi.responses import FileResponse
from utils.jobs import enqueue
from utils.read.custom_emoji import current_custom_emoji
from utils.storage import ContentAddressedStorage
router = APIRouter(
prefix="/api/custom-emoji", tags=["custom-emoji"], route_class=DishkaRoute
)
@router.get("/{custom_emoji_id}")
async def serve_custom_emoji(
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
custom_emoji_id: int,
account_id: Annotated[int, Query()],
) -> FileResponse:
emoji = await current_custom_emoji(pool, custom_emoji_id)
if emoji is None or not emoji.downloaded or emoji.storage_key is None:
await enqueue(
pool, account_id, "fetch_custom_emoji", {"custom_emoji_id": custom_emoji_id}
)
raise HTTPException(
status_code=409, detail="custom emoji not downloaded; fetching"
)
return FileResponse(
storage.url(emoji.storage_key),
media_type=emoji.mime or "application/octet-stream",
)
+119
View File
@@ -0,0 +1,119 @@
import asyncio
import json
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, Query
from pydantic import BaseModel
from api.routers.policy import POLICY_CHANGED_CHANNEL
from utils.jobs import enqueue
from utils.policy import repository as policy_repository
from utils.policy.defaults import TRACKING
from utils.policy.models import ScopeType
from utils.read import discover
from utils.read.models import DiscoverItem
router = APIRouter(prefix="/api", tags=["discover"], route_class=DishkaRoute)
DEFAULT_LIMIT = 30
REMOTE_TIMEOUT_SECONDS = 20.0
POLL_INTERVAL_SECONDS = 0.2
FINISHED = ("done", "failed", "canceled")
_CHAT_POLICY_ID = """
SELECT id FROM capture_policy
WHERE account_id = $1 AND scope_type = 'chat' AND scope_id = $2
"""
AccountId = Annotated[int, Query()]
class TrackRequest(BaseModel):
account_id: int
backfill: bool = True
class SyncContactsRequest(BaseModel):
account_id: int
async def _remote_ids(
pool: asyncpg.Pool, account_id: int, query: str, limit: int
) -> list[int]:
job_id = await enqueue(
pool, account_id, "search_peers", {"query": query, "limit": limit}
)
loop = asyncio.get_running_loop()
deadline = loop.time() + REMOTE_TIMEOUT_SECONDS
while loop.time() < deadline:
await asyncio.sleep(POLL_INTERVAL_SECONDS)
row = await pool.fetchrow(
"SELECT status, progress FROM jobs WHERE id = $1", job_id
)
if row is not None and row["status"] in FINISHED:
await pool.execute("DELETE FROM jobs WHERE id = $1", job_id)
return json.loads(row["progress"]).get("ids", [])
return []
@router.get("/discover")
async def discover_peers(
pool: FromDishka[asyncpg.Pool],
account_id: AccountId,
query: Annotated[str, Query()] = "",
remote: Annotated[bool, Query()] = False,
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
) -> list[DiscoverItem]:
if not query.strip():
return []
items = await discover.search(pool, account_id, query, limit)
if not remote:
return items
known = {item.chat_id for item in items}
ids = await _remote_ids(pool, account_id, query, limit)
extra = await discover.by_ids(
pool, account_id, [chat_id for chat_id in ids if chat_id not in known]
)
return [*items, *extra]
@router.get("/discover/{chat_id}")
async def discover_chat(
pool: FromDishka[asyncpg.Pool], chat_id: int, account_id: AccountId
) -> DiscoverItem:
return await discover.get_item(pool, account_id, chat_id)
@router.post("/chats/{chat_id}/track", status_code=201)
async def track_chat(
pool: FromDishka[asyncpg.Pool], chat_id: int, body: TrackRequest
) -> DiscoverItem:
kind = await discover.chat_kind(pool, body.account_id, chat_id)
toggles = TRACKING[kind]
policy_id = await pool.fetchval(_CHAT_POLICY_ID, body.account_id, chat_id)
if policy_id is None:
await policy_repository.create_policy(
pool, body.account_id, ScopeType.CHAT, chat_id, toggles
)
else:
await policy_repository.update_policy(pool, policy_id, toggles)
await pool.execute(f"NOTIFY {POLICY_CHANGED_CHANNEL}")
await enqueue(pool, body.account_id, "enrich_chat", {"chat_id": chat_id})
if body.backfill:
await enqueue(
pool,
body.account_id,
"backfill",
{"chat_id": chat_id, "media": True, "full": True},
)
return await discover.get_item(pool, body.account_id, chat_id)
@router.post("/contacts/sync", status_code=201)
async def sync_contacts(
pool: FromDishka[asyncpg.Pool], body: SyncContactsRequest
) -> dict[str, int]:
job_id = await enqueue(pool, body.account_id, "sync_contacts", {})
return {"job_id": job_id}
+42
View File
@@ -0,0 +1,42 @@
import asyncio
import json
from collections.abc import AsyncGenerator
from typing import Annotated
from fastapi import APIRouter, Query
from fastapi.responses import StreamingResponse
from api.realtime import Subscriber, hub
router = APIRouter(prefix="/api/events", tags=["events"])
HEARTBEAT_SECONDS = 15
AccountId = Annotated[int, Query()]
ChatId = Annotated[int | None, Query()]
async def _stream(sub: Subscriber) -> AsyncGenerator[str]:
try:
yield ": connected\n\n"
while True:
try:
frame = await asyncio.wait_for(
sub.queue.get(), timeout=HEARTBEAT_SECONDS
)
except TimeoutError:
yield ": keepalive\n\n"
continue
yield f"event: {frame['type']}\ndata: {json.dumps(frame)}\n\n"
finally:
hub.unsubscribe(sub)
@router.get("")
async def events(account_id: AccountId, chat_id: ChatId = None) -> StreamingResponse:
sub = hub.subscribe(account_id, chat_id)
return StreamingResponse(
_stream(sub),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
+106
View File
@@ -0,0 +1,106 @@
from datetime import UTC, datetime
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, Query, Request
from fastapi.responses import FileResponse, PlainTextResponse
from utils.files import (
content_disposition,
counts_as_download,
is_inline_mime,
resolve_mime,
)
from utils.read import shares
from utils.storage import ContentAddressedStorage
router = APIRouter(tags=["files"], route_class=DishkaRoute)
_GONE = "This link is no longer available.\n"
_MISSING = "Not found.\n"
_NO_STORE = {"Cache-Control": "private, no-store", "X-Robots-Tag": "noindex, nofollow"}
def _client_ip(request: Request) -> str | None:
forwarded = request.headers.get("x-forwarded-for")
if forwarded:
return forwarded.split(",")[0].strip()
return request.client.host if request.client else None
def _expired(row: asyncpg.Record) -> bool:
expires_at = row["expires_at"]
return expires_at is not None and expires_at <= datetime.now(UTC)
def _exhausted(row: asyncpg.Record) -> bool:
limit = row["max_downloads"]
return limit is not None and row["download_count"] >= limit
async def _serve(
request: Request,
pool: asyncpg.Pool,
storage: ContentAddressedStorage,
token: str,
*,
dl: bool,
) -> FileResponse | PlainTextResponse:
row = await shares.share_by_token(pool, token)
if row is None:
return PlainTextResponse(_MISSING, status_code=404, headers=_NO_STORE)
method = request.method
agent = request.headers.get("user-agent")
ip = _client_ip(request)
counts = counts_as_download(method, agent, request.headers.get("range"))
if row["revoked_at"] is not None or _expired(row) or _exhausted(row):
await shares.record_hit(pool, row["id"], method, ip, agent, counted=False)
return PlainTextResponse(_GONE, status_code=410, headers=_NO_STORE)
if not storage.exists(row["storage_key"]):
return PlainTextResponse(_MISSING, status_code=404, headers=_NO_STORE)
if counts and not await shares.consume_download(pool, row["id"]):
await shares.record_hit(pool, row["id"], method, ip, agent, counted=False)
return PlainTextResponse(_GONE, status_code=410, headers=_NO_STORE)
await shares.record_hit(pool, row["id"], method, ip, agent, counted=counts)
mime = resolve_mime(row["kind"], row["mime"], row["file_name"])
attachment = dl or not is_inline_mime(mime)
return FileResponse(
storage.url(row["storage_key"]),
media_type=mime,
headers={
**_NO_STORE,
"Content-Disposition": content_disposition(
row["file_name"], attachment=attachment
),
},
)
@router.api_route("/f/{token}", methods=["GET", "HEAD"], response_model=None)
async def serve_shared_file(
request: Request,
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
token: str,
dl: Annotated[bool, Query()] = False,
) -> FileResponse | PlainTextResponse:
return await _serve(request, pool, storage, token, dl=dl)
@router.api_route("/f/{token}/{name}", methods=["GET", "HEAD"], response_model=None)
async def serve_shared_file_named(
request: Request,
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
token: str,
name: str, # noqa: ARG001
dl: Annotated[bool, Query()] = False,
) -> FileResponse | PlainTextResponse:
return await _serve(request, pool, storage, token, dl=dl)
+72 -4
View File
@@ -1,14 +1,32 @@
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException
from fastapi import APIRouter, HTTPException, Query
from fastapi.responses import FileResponse
from utils.read.media import get_media
from utils.read.models import MediaView
from utils.cache import DAY_HEADERS, IMMUTABLE_HEADERS
from utils.files import content_disposition, media_file_name, resolve_mime
from utils.read.media import (
get_media,
get_media_version,
get_media_versions,
get_message_media,
)
from utils.read.models import MediaVersionView, MediaView
from utils.storage import ContentAddressedStorage
router = APIRouter(prefix="/api/media", tags=["media"], route_class=DishkaRoute)
Download = Annotated[bool, Query()]
def _attachment(headers: dict[str, str], file_name: str) -> dict[str, str]:
return {
**headers,
"Content-Disposition": content_disposition(file_name, attachment=True),
}
@router.get("/{media_id}/meta")
async def media_meta(pool: FromDishka[asyncpg.Pool], media_id: int) -> MediaView:
@@ -18,11 +36,57 @@ async def media_meta(pool: FromDishka[asyncpg.Pool], media_id: int) -> MediaView
return media
@router.get("/versions/{chat_id}/{message_id}")
async def message_media_versions(
pool: FromDishka[asyncpg.Pool],
chat_id: int,
message_id: int,
account_id: Annotated[int, Query()],
) -> list[MediaVersionView]:
return await get_media_versions(pool, account_id, chat_id, message_id)
@router.get("/version/{version_id}")
async def serve_media_version(
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
version_id: int,
download: Download = False,
) -> FileResponse:
version = await get_media_version(pool, version_id)
if version is None:
raise HTTPException(status_code=404, detail="media version not found")
headers = IMMUTABLE_HEADERS
if download:
headers = _attachment(
headers, media_file_name(version.kind, version.mime, version_id)
)
return FileResponse(
storage.url(version.storage_key),
media_type=resolve_mime(version.kind, version.mime),
headers=headers,
)
@router.get("/message/{chat_id}/{message_id}")
async def message_media(
pool: FromDishka[asyncpg.Pool],
chat_id: int,
message_id: int,
account_id: Annotated[int, Query()],
) -> MediaView:
media = await get_message_media(pool, account_id, chat_id, message_id)
if media is None:
raise HTTPException(status_code=404, detail="media not found")
return media
@router.get("/{media_id}")
async def serve_media(
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
media_id: int,
download: Download = False,
) -> FileResponse:
media = await get_media(pool, media_id)
if media is None:
@@ -32,7 +96,11 @@ async def serve_media(
status_code=409,
detail="media not downloaded; enqueue fetch via POST /api/media/fetch",
)
headers = DAY_HEADERS
if download:
headers = _attachment(headers, media.file_name or f"media_{media_id}")
return FileResponse(
storage.url(media.storage_key),
media_type=media.mime or "application/octet-stream",
media_type=resolve_mime(media.kind, media.mime, media.file_name),
headers=headers,
)
+9 -14
View File
@@ -5,13 +5,21 @@ from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, Query
from utils.read import peers
from utils.read.models import DEFAULT_LIMIT, Page, PeerHistoryView, PeerView, StoryView
from utils.read.models import PeerHistoryView, PeerView
router = APIRouter(prefix="/api", tags=["peers"], route_class=DishkaRoute)
AccountId = Annotated[int, Query()]
@router.get("/peers/batch")
async def get_peers(
pool: FromDishka[asyncpg.Pool], account_id: AccountId, ids: Annotated[str, Query()]
) -> list[PeerView]:
parsed = [int(part) for part in ids.split(",") if part.strip()]
return await peers.get_peers(pool, account_id, parsed)
@router.get("/peers/{peer_id}")
async def get_peer(
pool: FromDishka[asyncpg.Pool], peer_id: int, account_id: AccountId
@@ -27,16 +35,3 @@ async def peer_history(
pool: FromDishka[asyncpg.Pool], peer_id: int, account_id: AccountId
) -> list[PeerHistoryView]:
return await peers.get_peer_history(pool, account_id, peer_id)
@router.get("/stories")
async def stories(
pool: FromDishka[asyncpg.Pool],
account_id: AccountId,
peer_id: Annotated[int | None, Query()] = None,
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
offset: Annotated[int, Query()] = 0,
) -> list[StoryView]:
return await peers.get_stories(
pool, account_id, Page(limit=limit, offset=offset), peer_id=peer_id
)
+8 -2
View File
@@ -71,9 +71,15 @@ async def get_policy(pool: FromDishka[asyncpg.Pool], policy_id: int) -> PolicyRe
@router.put("/{policy_id}")
async def update_policy(
pool: FromDishka[asyncpg.Pool], policy_id: int, body: CaptureToggles
pool: FromDishka[asyncpg.Pool],
policy_id: int,
body: CaptureToggles,
account_id: Annotated[int | None, Query()] = None,
) -> PolicyRecord:
record = await repository.update_policy(pool, policy_id, body)
if account_id is None:
record = await repository.update_policy(pool, policy_id, body)
else:
record = await repository.override_policy(pool, policy_id, account_id, body)
if record is None:
raise HTTPException(status_code=404, detail="policy not found")
await pool.execute(f"NOTIFY {POLICY_CHANGED_CHANNEL}")
+7
View File
@@ -36,6 +36,13 @@ async def presence_history(
)
@router.get("/current")
async def current_presence(
pool: FromDishka[asyncpg.Pool], account_id: AccountId, peer_id: PeerId
) -> PresenceSample | None:
return await presence.current_presence(pool, account_id, peer_id)
@router.get("/hourly")
async def presence_hourly(
pool: FromDishka[asyncpg.Pool],
+70
View File
@@ -0,0 +1,70 @@
from datetime import datetime
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, Query
from utils.read import profile
from utils.read.models import (
DEFAULT_LIMIT,
ChatLinkView,
DayCount,
MediaView,
MessageAt,
Page,
)
router = APIRouter(
prefix="/api/chats/{chat_id}", tags=["profile"], route_class=DishkaRoute
)
AccountId = Annotated[int, Query()]
@router.get("/media")
async def chat_media(
pool: FromDishka[asyncpg.Pool],
chat_id: int,
account_id: AccountId,
kinds: Annotated[str, Query()],
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
offset: Annotated[int, Query()] = 0,
) -> list[MediaView]:
parsed = [part for part in kinds.split(",") if part.strip()]
return await profile.chat_media(
pool, account_id, chat_id, parsed, Page(limit=limit, offset=offset)
)
@router.get("/links")
async def chat_links(
pool: FromDishka[asyncpg.Pool],
chat_id: int,
account_id: AccountId,
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
offset: Annotated[int, Query()] = 0,
) -> list[ChatLinkView]:
return await profile.chat_links(
pool, account_id, chat_id, Page(limit=limit, offset=offset)
)
@router.get("/calendar")
async def chat_calendar(
pool: FromDishka[asyncpg.Pool], chat_id: int, account_id: AccountId
) -> list[DayCount]:
return await profile.daily_counts(pool, account_id, chat_id)
@router.get("/message-at")
async def message_at(
pool: FromDishka[asyncpg.Pool],
chat_id: int,
account_id: AccountId,
date: Annotated[datetime, Query()],
) -> MessageAt:
found = await profile.first_message_on_day(pool, account_id, chat_id, date)
if found is None:
raise HTTPException(status_code=404, detail="no message on day")
return found
+233
View File
@@ -0,0 +1,233 @@
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, Query
from pydantic import BaseModel
from utils.files import (
expiry_from_seconds,
media_file_name,
resolve_mime,
story_file_name,
)
from utils.read import chats, peers, shares
from utils.read.media import get_media, get_media_version
from utils.read.models import DEFAULT_LIMIT, FileShareHitView, FileShareView, Page
router = APIRouter(prefix="/api/shares", tags=["shares"], route_class=DishkaRoute)
class ShareCreate(BaseModel):
account_id: int
kind: str
media_id: int | None = None
version_id: int | None = None
peer_id: int | None = None
story_id: int | None = None
expires_in_seconds: int | None = None
max_downloads: int | None = None
class ShareUpdate(BaseModel):
expires_in_seconds: int | None = None
max_downloads: int | None = None
keep_expiry: bool = False
class Subject(BaseModel):
storage_key: str
file_name: str
mime: str | None
file_size: int | None
title: str | None
chat_id: int | None = None
message_id: int | None = None
peer_id: int | None = None
story_id: int | None = None
_NOT_STORED = "file is not downloaded yet"
async def _chat_title(pool: asyncpg.Pool, account_id: int, chat_id: int) -> str | None:
chat = await chats.get_chat(pool, account_id, chat_id)
return chat.title if chat else None
async def _media_subject(pool: asyncpg.Pool, media_id: int) -> Subject:
media = await get_media(pool, media_id)
if media is None:
raise HTTPException(status_code=404, detail="media not found")
if not media.downloaded or media.storage_key is None:
raise HTTPException(status_code=409, detail=_NOT_STORED)
file_name = media.file_name or f"media_{media.id}"
return Subject(
storage_key=media.storage_key,
file_name=file_name,
mime=resolve_mime(media.kind, media.mime, file_name),
file_size=media.file_size,
title=await _chat_title(pool, media.account_id, media.chat_id),
chat_id=media.chat_id,
message_id=media.message_id,
)
async def _version_subject(pool: asyncpg.Pool, version_id: int) -> Subject:
version = await get_media_version(pool, version_id)
if version is None:
raise HTTPException(status_code=404, detail="media version not found")
version_name = media_file_name(version.kind, version.mime, version_id)
return Subject(
storage_key=version.storage_key,
file_name=version_name,
mime=resolve_mime(version.kind, version.mime, version_name),
file_size=version.file_size,
title=None,
chat_id=None,
message_id=None,
)
async def _story_subject(
pool: asyncpg.Pool, account_id: int, peer_id: int, story_id: int
) -> Subject:
story = await peers.get_story(pool, account_id, peer_id, story_id)
if story is None:
raise HTTPException(status_code=404, detail="story not found")
if not story.downloaded or story.storage_key is None:
raise HTTPException(status_code=409, detail=_NOT_STORED)
return Subject(
storage_key=story.storage_key,
file_name=story_file_name(peer_id, story_id, story.media_kind),
mime="video/mp4" if story.media_kind == "video" else "image/jpeg",
file_size=None,
title=await _chat_title(pool, account_id, peer_id),
peer_id=peer_id,
story_id=story_id,
)
async def _resolve(pool: asyncpg.Pool, body: ShareCreate) -> Subject:
if body.kind == "media" and body.media_id is not None:
return await _media_subject(pool, body.media_id)
if body.kind == "media_version" and body.version_id is not None:
return await _version_subject(pool, body.version_id)
if body.kind == "story" and body.peer_id is not None and body.story_id is not None:
return await _story_subject(pool, body.account_id, body.peer_id, body.story_id)
raise HTTPException(status_code=422, detail="unsupported share subject")
@router.get("")
async def list_shares(
pool: FromDishka[asyncpg.Pool],
account_id: Annotated[int, Query()],
active_only: Annotated[bool, Query()] = False,
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
offset: Annotated[int, Query()] = 0,
) -> list[FileShareView]:
return await shares.list_shares(
pool, account_id, Page(limit=limit, offset=offset), active_only=active_only
)
@router.post("", status_code=201)
async def create_share(
pool: FromDishka[asyncpg.Pool], body: ShareCreate
) -> FileShareView:
subject = await _resolve(pool, body)
existing = await shares.find_active_share(
pool, body.account_id, subject.storage_key
)
if existing is not None:
return existing
return await shares.create_share(
pool,
body.account_id,
body.kind,
subject.storage_key,
subject.file_name,
mime=subject.mime,
file_size=subject.file_size,
title=subject.title,
chat_id=subject.chat_id,
message_id=subject.message_id,
peer_id=subject.peer_id,
story_id=subject.story_id,
expires_at=expiry_from_seconds(body.expires_in_seconds),
max_downloads=body.max_downloads,
)
@router.get("/lookup")
async def lookup_share(
pool: FromDishka[asyncpg.Pool],
account_id: Annotated[int, Query()],
kind: Annotated[str, Query()],
media_id: Annotated[int | None, Query()] = None,
version_id: Annotated[int | None, Query()] = None,
peer_id: Annotated[int | None, Query()] = None,
story_id: Annotated[int | None, Query()] = None,
) -> FileShareView | None:
body = ShareCreate(
account_id=account_id,
kind=kind,
media_id=media_id,
version_id=version_id,
peer_id=peer_id,
story_id=story_id,
)
subject = await _resolve(pool, body)
return await shares.find_active_share(pool, account_id, subject.storage_key)
@router.get("/{share_id}/hits")
async def share_hits(
pool: FromDishka[asyncpg.Pool],
share_id: int,
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
offset: Annotated[int, Query()] = 0,
) -> list[FileShareHitView]:
return await shares.list_hits(pool, share_id, Page(limit=limit, offset=offset))
@router.patch("/{share_id}")
async def update_share(
pool: FromDishka[asyncpg.Pool], share_id: int, body: ShareUpdate
) -> FileShareView:
current = await shares.get_share(pool, share_id)
if current is None:
raise HTTPException(status_code=404, detail="share not found")
expires_at = (
current.expires_at
if body.keep_expiry
else expiry_from_seconds(body.expires_in_seconds)
)
updated = await shares.update_share(
pool, share_id, expires_at=expires_at, max_downloads=body.max_downloads
)
if updated is None:
raise HTTPException(status_code=404, detail="share not found")
return updated
@router.post("/{share_id}/revoke")
async def revoke_share(pool: FromDishka[asyncpg.Pool], share_id: int) -> FileShareView:
share = await shares.revoke_share(pool, share_id)
if share is None:
raise HTTPException(status_code=404, detail="share not found")
return share
@router.post("/{share_id}/reissue")
async def reissue_share(pool: FromDishka[asyncpg.Pool], share_id: int) -> FileShareView:
share = await shares.rotate_token(pool, share_id)
if share is None:
raise HTTPException(status_code=404, detail="share not found")
return share
@router.delete("/{share_id}", status_code=204)
async def delete_share(pool: FromDishka[asyncpg.Pool], share_id: int) -> None:
if not await shares.delete_share(pool, share_id):
raise HTTPException(status_code=404, detail="share not found")
+56
View File
@@ -0,0 +1,56 @@
from typing import Annotated
import asyncpg
from dishka.integrations.fastapi import DishkaRoute, FromDishka
from fastapi import APIRouter, HTTPException, Query
from fastapi.responses import FileResponse
from utils.files import content_disposition, story_file_name
from utils.read import peers
from utils.read.models import DEFAULT_LIMIT, Page, StoryView
from utils.storage import ContentAddressedStorage
router = APIRouter(prefix="/api", tags=["stories"], route_class=DishkaRoute)
AccountId = Annotated[int, Query()]
_STORY_MIME = {"photo": "image/jpeg", "video": "video/mp4"}
@router.get("/stories")
async def list_stories(
pool: FromDishka[asyncpg.Pool],
account_id: AccountId,
peer_id: Annotated[int | None, Query()] = None,
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
offset: Annotated[int, Query()] = 0,
) -> list[StoryView]:
return await peers.get_stories(
pool, account_id, Page(limit=limit, offset=offset), peer_id=peer_id
)
@router.get("/stories/{peer_id}/{story_id}/media")
async def serve_story_media(
pool: FromDishka[asyncpg.Pool],
storage: FromDishka[ContentAddressedStorage],
peer_id: int,
story_id: int,
account_id: AccountId,
download: Annotated[bool, Query()] = False,
) -> FileResponse:
story = await peers.get_story(pool, account_id, peer_id, story_id)
if story is None:
raise HTTPException(status_code=404, detail="story not found")
if not story.downloaded or story.storage_key is None:
raise HTTPException(status_code=409, detail="story media not downloaded")
headers = {}
if download:
headers["Content-Disposition"] = content_disposition(
story_file_name(peer_id, story_id, story.media_kind), attachment=True
)
return FileResponse(
storage.url(story.storage_key),
media_type=_STORY_MIME.get(story.media_kind or "", "application/octet-stream"),
headers=headers,
)
+2 -2
View File
@@ -1,3 +1,3 @@
from userbot.modules.client import PyroClient
from userbot.modules.client import DEVICE_MODEL, DEVICE_MODEL_LIMIT, PyroClient
__all__ = ["PyroClient"]
__all__ = ["DEVICE_MODEL", "DEVICE_MODEL_LIMIT", "PyroClient"]
+5
View File
@@ -3,6 +3,7 @@ from pyrogram.types import Message
from userbot import PyroClient
from userbot.modules.capture import repository
from userbot.modules.capture.repository import CHANNEL_ID_THRESHOLD
from utils.events import notify_bg_event
@PyroClient.on_deleted_messages()
@@ -21,8 +22,12 @@ async def on_deleted_messages(client: PyroClient, messages: list[Message]) -> No
channels.setdefault(chat_id, []).append(message.id)
if box:
await repository.mark_deleted_box(ctx.pool, ctx.account_id, box)
await notify_bg_event(ctx.pool, "delete", ctx.account_id, message_ids=box)
for chat_id, ids in channels.items():
await repository.mark_deleted_channel(ctx.pool, ctx.account_id, chat_id, ids)
await notify_bg_event(
ctx.pool, "delete", ctx.account_id, chat_id=chat_id, message_ids=ids
)
handlers = on_deleted_messages.handlers
+18 -3
View File
@@ -4,13 +4,20 @@ from userbot import PyroClient
from userbot.modules.capture import repository
from userbot.modules.capture.chat_meta import meta_from_chat
from userbot.modules.capture.message import sender_id
from userbot.modules.media import self_destruct_ttl
from userbot.modules.media import capture_media, media_unique_id, self_destruct_ttl
from utils.events import notify_bg_event
@PyroClient.on_edited_message()
async def on_edited_message(client: PyroClient, message: Message) -> None:
ctx = client.capture
if ctx is None or message.empty or message.chat is None or message.date is None:
if (
ctx is None
or message.empty
or message.scheduled
or message.chat is None
or message.date is None
):
return
chat = message.chat
chat_id = chat.id or 0
@@ -18,7 +25,7 @@ async def on_edited_message(client: PyroClient, message: Message) -> None:
toggles = ctx.resolve(meta)
if not toggles.track_edits_deletes:
return
await repository.add_version(
changed = await repository.add_version(
ctx.pool,
ctx.account_id,
chat_id,
@@ -28,9 +35,17 @@ async def on_edited_message(client: PyroClient, message: Message) -> None:
message.text or message.caption,
str(message),
message.edit_date,
media_unique_id(message),
has_media=message.media is not None,
is_self_destruct=self_destruct_ttl(message) is not None,
)
if not changed:
return
if message.media is not None:
await capture_media(client, message, ctx, chat_id, message.id, toggles)
await notify_bg_event(
ctx.pool, "edit", ctx.account_id, chat_id=chat_id, message_id=message.id
)
handlers = on_edited_message.handlers
+11 -1
View File
@@ -5,12 +5,19 @@ from userbot.modules.capture import capture_message
from userbot.modules.capture.chat_meta import meta_from_chat
from userbot.modules.stt import is_transcribable
from userbot.modules.stt.gate import safe_transcribe
from utils.events import notify_bg_event
@PyroClient.on_message()
async def on_message(client: PyroClient, message: Message) -> None:
ctx = client.capture
if ctx is None or message.empty or message.chat is None or message.date is None:
if (
ctx is None
or message.empty
or message.scheduled
or message.chat is None
or message.date is None
):
return
meta = meta_from_chat(message.chat, ctx.contacts.ids)
await ctx.watches.on_text(meta.chat_id, message.id, message.text or message.caption)
@@ -18,6 +25,9 @@ async def on_message(client: PyroClient, message: Message) -> None:
if not toggles.messages:
return
await capture_message(client, message, ctx, toggles)
await notify_bg_event(
ctx.pool, "message", ctx.account_id, chat_id=meta.chat_id, message_id=message.id
)
if (
toggles.stt
and is_transcribable(message)
+2
View File
@@ -4,6 +4,7 @@ from pyrogram.types import User
from userbot import PyroClient
from userbot.modules.presence import repository
from utils.events import notify_bg_event
@PyroClient.on_user_status()
@@ -22,6 +23,7 @@ async def on_user_status(client: PyroClient, user: User) -> None:
str(user.raw),
)
await ctx.watches.on_status(user.id, is_online=user.status.name.lower() == "online")
await notify_bg_event(ctx.pool, "presence", ctx.account_id, chat_id=user.id)
handlers = on_user_status.handlers
@@ -3,6 +3,7 @@ from pyrogram import raw, utils
from userbot import PyroClient
from userbot.modules.capture import repository
from userbot.modules.capture.chat_meta import meta_from_peer
from utils.events import notify_bg_event
HANDLES = (raw.types.UpdateMessageReactions,)
@@ -47,3 +48,10 @@ async def handle(
await repository.sync_reactions(
ctx.pool, ctx.account_id, meta.chat_id, update.msg_id, current
)
await notify_bg_event(
ctx.pool,
"reaction",
ctx.account_id,
chat_id=meta.chat_id,
message_id=update.msg_id,
)
@@ -4,6 +4,7 @@ from pyrogram import raw, utils
from userbot import PyroClient
from userbot.modules.read_receipts import repository
from utils.events import notify_bg_event
HANDLES = (raw.types.UpdateReadHistoryOutbox,)
@@ -24,3 +25,6 @@ async def handle(
update.max_id,
str(update),
)
await notify_bg_event(
ctx.pool, "receipt", ctx.account_id, chat_id=chat_id, message_id=update.max_id
)
+3 -41
View File
@@ -1,52 +1,14 @@
from io import BytesIO
from pyrogram.types import Story
from userbot import PyroClient
from userbot.modules.stories import repository
def _peer_id(story: Story) -> int:
if story.chat is not None:
return story.chat.id or 0
if story.from_user is not None:
return story.from_user.id or 0
return 0
from userbot.modules.stories.service import save_story
@PyroClient.on_story()
async def on_story(client: PyroClient, story: Story) -> None:
ctx = client.capture
if ctx is None:
if client.capture is None:
return
media_kind = story.media.name.lower() if story.media else None
storage_key: str | None = None
file_size: int | None = None
downloaded = False
if not story.deleted and story.media is not None:
buffer = await client.download_media(story, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
storage_key = ctx.storage.put(data)
file_size = len(data)
downloaded = True
await repository.upsert_story(
ctx.pool,
ctx.account_id,
_peer_id(story),
story.id,
story.date,
story.expire_date,
story.caption,
media_kind,
storage_key,
file_size,
story.views,
str(story.raw),
pinned=bool(story.pinned),
deleted=bool(story.deleted),
downloaded=downloaded,
)
await save_story(client, client.capture, story)
handlers = on_story.handlers
@@ -1,3 +1,4 @@
from userbot.modules.avatars.downloader import capture_avatar
from userbot.modules.avatars.repository import note_avatar
__all__ = ["capture_avatar"]
__all__ = ["capture_avatar", "note_avatar"]
@@ -1,3 +1,5 @@
import json
import asyncpg
_INSERT_AVATAR = """
@@ -13,6 +15,16 @@ SELECT 1 FROM avatars
WHERE account_id = $1 AND owner_id = $2 AND unique_id = $3
"""
_GET_FILE = """
SELECT raw ->> 'file_id' AS file_id, downloaded FROM avatars
WHERE account_id = $1 AND owner_kind = $2 AND owner_id = $3 AND unique_id = $4
"""
_MARK_DOWNLOADED = """
UPDATE avatars SET downloaded = true, storage_key = $5, file_size = $6
WHERE account_id = $1 AND owner_kind = $2 AND owner_id = $3 AND unique_id = $4
"""
async def avatar_exists(
pool: asyncpg.Pool, account_id: int, owner_id: int, unique_id: str
@@ -46,3 +58,54 @@ async def insert_avatar( # noqa: PLR0913
downloaded,
raw,
)
async def note_avatar( # noqa: PLR0913
pool: asyncpg.Pool,
account_id: int,
owner_id: int,
owner_kind: str,
unique_id: str,
file_id: str,
) -> None:
await insert_avatar(
pool,
account_id,
owner_id,
owner_kind,
unique_id,
None,
None,
None,
json.dumps({"file_id": file_id}),
downloaded=False,
)
async def get_avatar_file(
pool: asyncpg.Pool, account_id: int, owner_kind: str, owner_id: int, unique_id: str
) -> tuple[str | None, bool] | None:
row = await pool.fetchrow(_GET_FILE, account_id, owner_kind, owner_id, unique_id)
if row is None:
return None
return row["file_id"], row["downloaded"]
async def mark_avatar_downloaded( # noqa: PLR0913
pool: asyncpg.Pool,
account_id: int,
owner_kind: str,
owner_id: int,
unique_id: str,
storage_key: str,
file_size: int,
) -> None:
await pool.execute(
_MARK_DOWNLOADED,
account_id,
owner_kind,
owner_id,
unique_id,
storage_key,
file_size,
)
@@ -1,6 +1,7 @@
import asyncpg
from pyrogram import Client
from userbot.modules.capture.identity import ChatMetaCache, PeerIdentityCache
from userbot.modules.contacts import ContactCache
from userbot.modules.folders import FolderCache
from userbot.modules.watches import WatchCache
@@ -25,6 +26,8 @@ class CaptureContext:
self.folders = folders
self.contacts = contacts
self.watches = WatchCache(pool, account_id)
self.peer_identity = PeerIdentityCache()
self.chat_meta = ChatMetaCache()
self.policies = PolicySet()
async def reload_policies(self) -> None:
@@ -0,0 +1,110 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from userbot.modules.capture.chat_meta import chat_kind
from userbot.modules.profiles.parse import ProfileFields, snapshot_from_high_level
from userbot.modules.profiles.repository import get_peer, write_profile
from utils.policy.models import ChatKind
if TYPE_CHECKING:
import asyncpg
from pyrogram.types import Message
from userbot.modules.capture.context import CaptureContext
class PeerIdentityCache:
def __init__(self) -> None:
self._cache: dict[int, ProfileFields | None] = {}
async def changed(
self, pool: asyncpg.Pool, account_id: int, peer_id: int, fields: ProfileFields
) -> bool:
if peer_id not in self._cache:
self._cache[peer_id] = await get_peer(pool, account_id, peer_id)
if self._cache[peer_id] == fields:
return False
self._cache[peer_id] = fields
return True
class ChatMetaCache:
def __init__(self) -> None:
self._cache: dict[int, tuple[str | None, str | None]] = {}
async def changed(
self,
pool: asyncpg.Pool,
account_id: int,
chat_id: int,
meta: tuple[str | None, str | None],
) -> bool:
from userbot.modules.groups.repository import ( # noqa: PLC0415
get_latest_chat_meta,
)
if chat_id not in self._cache:
self._cache[chat_id] = await get_latest_chat_meta(pool, account_id, chat_id)
if self._cache[chat_id] == meta:
return False
self._cache[chat_id] = meta
return True
async def _capture_peer(message: Message, ctx: CaptureContext) -> None:
user = message.from_user
if user is None:
return
fields, photo_file_id, photo_unique_id = snapshot_from_high_level(user)
if not await ctx.peer_identity.changed(ctx.pool, ctx.account_id, user.id, fields):
return
await write_profile(ctx.pool, ctx.account_id, user.id, fields, str(user))
if photo_file_id and photo_unique_id:
from userbot.modules.avatars import note_avatar # noqa: PLC0415
await note_avatar(
ctx.pool, ctx.account_id, user.id, "peer", photo_unique_id, photo_file_id
)
async def _capture_chat(message: Message, ctx: CaptureContext) -> None:
chat = message.chat
if (
chat is None
or chat.id is None
or message.date is None
or chat_kind(chat.type) is ChatKind.DM
):
return
photo = chat.photo
photo_unique_id = photo.big_photo_unique_id if photo else None
photo_file_id = photo.big_file_id if photo else None
meta = (chat.title, photo_unique_id)
if not await ctx.chat_meta.changed(ctx.pool, ctx.account_id, chat.id, meta):
return
from userbot.modules.groups.repository import insert_chat_history # noqa: PLC0415
await insert_chat_history(
ctx.pool,
ctx.account_id,
chat.id,
message.id,
"meta",
chat.title,
photo_unique_id,
None,
message.date,
str(message),
)
if photo_file_id and photo_unique_id:
from userbot.modules.avatars import note_avatar # noqa: PLC0415
await note_avatar(
ctx.pool, ctx.account_id, chat.id, "chat", photo_unique_id, photo_file_id
)
async def capture_identity(message: Message, ctx: CaptureContext) -> None:
await _capture_peer(message, ctx)
await _capture_chat(message, ctx)
@@ -51,6 +51,9 @@ async def capture_message(
has_media=message.media is not None,
is_self_destruct=self_destruct_ttl(message) is not None,
)
from userbot.modules.capture.identity import capture_identity # noqa: PLC0415
await capture_identity(message, ctx)
await capture_media(client, message, ctx, chat_id, message.id, toggles)
buttons = callbacks(message)
if buttons:
+102 -17
View File
@@ -23,28 +23,71 @@ INSERT INTO messages
has_media, is_self_destruct, edited_at)
VALUES ($1, $2, $3, $4, $5, $6, $7::jsonb, $8, $9, now())
ON CONFLICT (account_id, chat_id, message_id, date) DO UPDATE SET
text = EXCLUDED.text,
raw = EXCLUDED.raw,
has_media = EXCLUDED.has_media,
is_self_destruct = EXCLUDED.is_self_destruct,
edited_at = now()
"""
_INSERT_VERSION = """
INSERT INTO message_versions
(account_id, chat_id, message_id, observed_at, edit_date, text, raw)
VALUES ($1, $2, $3, now(), $4, $5, $6::jsonb)
VALUES ($1, $2, $3, clock_timestamp(), $4, $5, $6::jsonb)
ON CONFLICT DO NOTHING
"""
_SNAPSHOT_ORIGINAL = """
INSERT INTO message_versions
(account_id, chat_id, message_id, observed_at, edit_date, text, raw)
SELECT account_id, chat_id, message_id, clock_timestamp(), NULL, text, raw
FROM messages m
WHERE m.account_id = $1 AND m.chat_id = $2 AND m.message_id = $3
AND NOT EXISTS (
SELECT 1 FROM message_versions v
WHERE v.account_id = m.account_id AND v.chat_id = m.chat_id
AND v.message_id = m.message_id
)
ON CONFLICT DO NOTHING
"""
_CURRENT_CONTENT = """
SELECT
m.text AS text,
(SELECT d.unique_id FROM media d
WHERE d.account_id = m.account_id AND d.chat_id = m.chat_id
AND d.message_id = m.message_id) AS media_unique_id
FROM messages m
WHERE m.account_id = $1 AND m.chat_id = $2 AND m.message_id = $3
ORDER BY m.date DESC LIMIT 1
"""
_CURRENT_MEDIA = """
SELECT unique_id, storage_key, file_size, downloaded FROM media
WHERE account_id = $1 AND chat_id = $2 AND message_id = $3
"""
_INSERT_MEDIA_VERSION = """
INSERT INTO media_versions
(account_id, chat_id, message_id, observed_at, kind, storage_key,
file_size, mime, ttl_seconds, unique_id)
VALUES ($1, $2, $3, clock_timestamp(), $4, $5, $6, $7, $8, $9)
ON CONFLICT (account_id, chat_id, message_id, storage_key) DO NOTHING
"""
_INSERT_MEDIA = """
INSERT INTO media
(account_id, chat_id, message_id, kind, storage_key, file_size, mime,
ttl_seconds, downloaded)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)
ttl_seconds, downloaded, unique_id)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
ON CONFLICT (account_id, chat_id, message_id) DO UPDATE SET
kind = EXCLUDED.kind,
storage_key = EXCLUDED.storage_key,
file_size = EXCLUDED.file_size,
mime = EXCLUDED.mime,
ttl_seconds = EXCLUDED.ttl_seconds,
downloaded = EXCLUDED.downloaded
downloaded = EXCLUDED.downloaded,
unique_id = EXCLUDED.unique_id
"""
@@ -75,6 +118,16 @@ async def upsert_message( # noqa: PLR0913
)
async def max_message_id(
pool: asyncpg.Pool, account_id: int, chat_id: int
) -> int | None:
return await pool.fetchval(
"SELECT max(message_id) FROM messages WHERE account_id = $1 AND chat_id = $2",
account_id,
chat_id,
)
async def mark_deleted_box(
pool: asyncpg.Pool, account_id: int, message_ids: list[int]
) -> None:
@@ -111,11 +164,20 @@ async def add_version( # noqa: PLR0913
text: str | None,
raw: str,
edit_date: datetime | None,
media_unique_id: str | None,
*,
has_media: bool,
is_self_destruct: bool,
) -> None:
) -> bool:
async with pool.acquire() as conn, conn.transaction():
current = await conn.fetchrow(_CURRENT_CONTENT, account_id, chat_id, message_id)
text_changed = current is None or current["text"] != text
media_changed = (
current is not None and current["media_unique_id"] != media_unique_id
)
if not (text_changed or media_changed):
return False
await conn.execute(_SNAPSHOT_ORIGINAL, account_id, chat_id, message_id)
await conn.execute(
_TOUCH_EDITED,
account_id,
@@ -131,6 +193,13 @@ async def add_version( # noqa: PLR0913
await conn.execute(
_INSERT_VERSION, account_id, chat_id, message_id, edit_date, text, raw
)
return True
async def current_media(
pool: asyncpg.Pool, account_id: int, chat_id: int, message_id: int
) -> asyncpg.Record | None:
return await pool.fetchrow(_CURRENT_MEDIA, account_id, chat_id, message_id)
async def insert_media( # noqa: PLR0913
@@ -143,21 +212,37 @@ async def insert_media( # noqa: PLR0913
file_size: int | None,
mime: str | None,
ttl_seconds: int | None,
unique_id: str | None,
*,
downloaded: bool,
) -> None:
await pool.execute(
_INSERT_MEDIA,
account_id,
chat_id,
message_id,
kind,
storage_key,
file_size,
mime,
ttl_seconds,
downloaded,
)
async with pool.acquire() as conn, conn.transaction():
await conn.execute(
_INSERT_MEDIA,
account_id,
chat_id,
message_id,
kind,
storage_key,
file_size,
mime,
ttl_seconds,
downloaded,
unique_id,
)
if storage_key is not None:
await conn.execute(
_INSERT_MEDIA_VERSION,
account_id,
chat_id,
message_id,
kind,
storage_key,
file_size,
mime,
ttl_seconds,
unique_id,
)
async def insert_callbacks(
+15 -4
View File
@@ -6,19 +6,30 @@ if TYPE_CHECKING:
from userbot.modules.capture import CaptureContext
DEVICE_MODEL = "Beavergram"
DEVICE_MODEL_LIMIT = 32
class PyroClient(Client):
def __init__(
self, name: str, *, workdir: str = "sessions", load_handlers: bool = True
self,
name: str,
*,
workdir: str = "sessions",
device_model: str | None = None,
load_handlers: bool = True,
) -> None:
super().__init__(
name,
workdir=workdir,
api_id=2040,
api_hash="b18441a1ff607e10a989891a5462e627",
device_model="Desktop",
device_model=device_model or DEVICE_MODEL,
system_version="Windows 11 x64",
app_version="6.2.4 x64",
app_version="7.0.8 x64",
lang_pack="tdesktop",
lang_code="en",
system_lang_code="en-US",
client_platform=enums.ClientPlatform.DESKTOP,
)
self.capture: CaptureContext | None = None
@@ -30,4 +41,4 @@ class PyroClient(Client):
self.add_handler(*handler)
__all__ = ["PyroClient"]
__all__ = ["DEVICE_MODEL", "DEVICE_MODEL_LIMIT", "PyroClient"]
@@ -0,0 +1,34 @@
import asyncpg
_GET = """
SELECT downloaded FROM custom_emoji WHERE custom_emoji_id = $1
"""
_UPSERT_DOWNLOADED = """
INSERT INTO custom_emoji
(custom_emoji_id, storage_key, file_size, mime, kind, downloaded, raw)
VALUES ($1, $2, $3, $4, $5, true, '{}'::jsonb)
ON CONFLICT (custom_emoji_id) DO UPDATE SET
storage_key = EXCLUDED.storage_key,
file_size = EXCLUDED.file_size,
mime = EXCLUDED.mime,
kind = EXCLUDED.kind,
downloaded = true
"""
async def is_downloaded(pool: asyncpg.Pool, custom_emoji_id: int) -> bool:
return bool(await pool.fetchval(_GET, custom_emoji_id))
async def upsert_downloaded( # noqa: PLR0913
pool: asyncpg.Pool,
custom_emoji_id: int,
storage_key: str,
file_size: int | None,
mime: str | None,
kind: str,
) -> None:
await pool.execute(
_UPSERT_DOWNLOADED, custom_emoji_id, storage_key, file_size, mime, kind
)
@@ -17,6 +17,26 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8::jsonb)
ON CONFLICT (account_id, chat_id, message_id, user_id) DO NOTHING
"""
_LATEST_TITLE = """
SELECT title FROM chat_history
WHERE account_id = $1 AND chat_id = $2 AND title IS NOT NULL
ORDER BY ts DESC LIMIT 1
"""
_LATEST_PHOTO = """
SELECT photo_unique_id FROM chat_history
WHERE account_id = $1 AND chat_id = $2 AND photo_unique_id IS NOT NULL
ORDER BY ts DESC LIMIT 1
"""
async def get_latest_chat_meta(
pool: asyncpg.Pool, account_id: int, chat_id: int
) -> tuple[str | None, str | None]:
title = await pool.fetchval(_LATEST_TITLE, account_id, chat_id)
photo_unique_id = await pool.fetchval(_LATEST_PHOTO, account_id, chat_id)
return title, photo_unique_id
async def insert_chat_history( # noqa: PLR0913
pool: asyncpg.Pool,
@@ -24,3 +24,6 @@ class JobContext:
async def report_progress(self, progress: dict[str, Any]) -> None:
self.job.progress = progress
await repository.report_progress(self.pool, self.job_id, progress)
async def is_canceled(self) -> bool:
return await repository.is_canceled(self.pool, self.job_id)
@@ -1,3 +1,25 @@
from userbot.modules.jobs.handlers import backfill, fetch_media, transcribe
from userbot.modules.jobs.handlers import (
backfill,
backfill_stories,
enrich_chat,
fetch_avatar,
fetch_custom_emoji,
fetch_media,
search_peers,
sync_contacts,
sync_dialogs,
transcribe,
)
__all__ = ["backfill", "fetch_media", "transcribe"]
__all__ = [
"backfill",
"backfill_stories",
"enrich_chat",
"fetch_avatar",
"fetch_custom_emoji",
"fetch_media",
"search_peers",
"sync_contacts",
"sync_dialogs",
"transcribe",
]
@@ -1,11 +1,48 @@
from pyrogram import Client
from pyrogram.errors import PeerIdInvalid
from pyrogram.types import Message
from userbot.modules.capture import capture_message
from userbot.modules.capture import repository as capture_repo
from userbot.modules.capture.chat_meta import meta_from_chat
from userbot.modules.capture.context import CaptureContext
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
from userbot.modules.stt import repository as stt_repo
from userbot.modules.stt import should_transcribe_on_backfill
from userbot.modules.stt.gate import safe_transcribe
from utils.policy.models import CaptureToggles
SAVE_EVERY = 100
async def resolve_min_id(ctx: JobContext, chat_id: int) -> int:
cursor = ctx.job.cursor or {}
if "min_id" in cursor:
return int(cursor["min_id"])
if ctx.job.params.get("full"):
return 0
newest = await capture_repo.max_message_id(ctx.pool, ctx.account_id, chat_id)
return newest + 1 if newest else 0
async def maybe_transcribe(
client: Client,
capture: CaptureContext,
chat_id: int,
message: Message,
self_id: int | None,
) -> None:
if not (should_transcribe_on_backfill(message, self_id) and message.chat):
return
meta = meta_from_chat(message.chat, capture.contacts.ids)
already = await stt_repo.is_transcribed(
capture.pool, capture.account_id, chat_id, message.id
)
if capture.resolve(meta).stt and not already:
await safe_transcribe(client, capture, chat_id, message.id)
@register("backfill")
async def backfill(ctx: JobContext) -> None:
client = ctx.client
@@ -20,14 +57,25 @@ async def backfill(ctx: JobContext) -> None:
media=bool(ctx.job.params.get("media")),
self_destruct_media=False,
)
max_id = (ctx.job.cursor or {}).get("max_id", 0)
max_id = int((ctx.job.cursor or {}).get("max_id", 0))
min_id = await resolve_min_id(ctx, chat_id)
await ctx.save_cursor({"max_id": max_id, "min_id": min_id})
processed = ctx.job.progress.get("processed", 0)
kwargs = {"max_id": max_id} if max_id else {}
async for message in client.get_chat_history(chat_id, **kwargs):
await capture_message(client, message, capture, toggles)
processed += 1
if processed % SAVE_EVERY == 0:
next_max = message.id - 1
await ctx.save_cursor({"max_id": next_max})
await ctx.report_progress({"processed": processed, "max_id": next_max})
self_id = client.me.id if client.me else None
try:
async for message in client.get_chat_history(
chat_id, max_id=max_id, min_id=min_id
):
await capture_message(client, message, capture, toggles)
await maybe_transcribe(client, capture, chat_id, message, self_id)
processed += 1
if processed % SAVE_EVERY == 0:
next_max = message.id - 1
await ctx.save_cursor({"max_id": next_max, "min_id": min_id})
await ctx.report_progress({"processed": processed, "max_id": next_max})
if await ctx.is_canceled():
return
except PeerIdInvalid:
await ctx.report_progress({"processed": processed, "error": "peer_id_invalid"})
return
await ctx.report_progress({"processed": processed, "done": True})
@@ -0,0 +1,70 @@
from collections.abc import AsyncIterator, Callable
from pyrogram import Client
from pyrogram.errors import FloodPremiumWait, FloodWait, RPCError
from pyrogram.types import Story
from userbot.modules.capture.context import CaptureContext
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
from userbot.modules.stories.service import save_story
SAVE_EVERY = 10
StorySource = Callable[[], AsyncIterator[Story]]
def _sources(client: Client, peer_id: int, *, own: bool) -> dict[str, StorySource]:
sources: dict[str, StorySource] = {
"active": lambda: client.get_chat_stories(peer_id),
"pinned": lambda: client.get_pinned_stories(peer_id),
}
if own:
sources["archived"] = lambda: client.get_archived_stories(peer_id)
return sources
async def _drain(
ctx: JobContext, capture: CaptureContext, name: str, source: StorySource
) -> int:
client = ctx.client
if client is None:
return 0
saved = 0
async for story in source():
try:
await save_story(client, capture, story)
except (FloodWait, FloodPremiumWait):
raise
except RPCError:
continue
saved += 1
if saved % SAVE_EVERY == 0:
await ctx.report_progress({"saved": saved, "source": name})
if await ctx.is_canceled():
break
return saved
@register("backfill_stories")
async def backfill_stories(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
peer_id = ctx.job.params["peer_id"]
own = client.me is not None and client.me.id == peer_id
saved = 0
errors: dict[str, str] = {}
for name, source in _sources(client, peer_id, own=own).items():
try:
saved += await _drain(ctx, capture, name, source)
except (FloodWait, FloodPremiumWait):
raise
except RPCError as exc:
errors[name] = type(exc).__name__
if await ctx.is_canceled():
break
await ctx.report_progress({"saved": saved, "done": True, "errors": errors})
@@ -0,0 +1,91 @@
from datetime import UTC, datetime
from pyrogram import Client
from pyrogram.errors import BadRequest, Forbidden
from pyrogram.types import User
from userbot.modules.avatars import note_avatar
from userbot.modules.capture.context import CaptureContext
from userbot.modules.groups.repository import insert_chat_history
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
from userbot.modules.profiles.parse import snapshot_from_high_level
from userbot.modules.profiles.repository import write_profile
MEMBER_CAP = 200
_MISSING_SENDERS = """
SELECT DISTINCT sender_id FROM messages
WHERE account_id = $1 AND chat_id = $2 AND sender_id > 0
AND NOT EXISTS (
SELECT 1 FROM peers p WHERE p.account_id = $1 AND p.peer_id = messages.sender_id
)
LIMIT 100
"""
async def _save_user(ctx: CaptureContext, user: User) -> None:
fields, photo_file_id, photo_unique_id = snapshot_from_high_level(user)
await write_profile(ctx.pool, ctx.account_id, user.id, fields, str(user))
if photo_file_id and photo_unique_id:
await note_avatar(
ctx.pool, ctx.account_id, user.id, "peer", photo_unique_id, photo_file_id
)
async def _enrich_chat_meta(client: Client, ctx: CaptureContext, chat_id: int) -> None:
chat = await client.get_chat(chat_id)
photo = chat.photo
photo_unique_id = photo.big_photo_unique_id if photo else None
photo_file_id = photo.big_file_id if photo else None
await insert_chat_history(
ctx.pool,
ctx.account_id,
chat_id,
0,
"meta",
chat.title,
photo_unique_id,
None,
datetime.now(UTC),
str(chat),
)
if photo_file_id and photo_unique_id:
await note_avatar(
ctx.pool, ctx.account_id, chat_id, "chat", photo_unique_id, photo_file_id
)
async def _enrich_members(client: Client, ctx: CaptureContext, chat_id: int) -> None:
try:
async for member in client.get_chat_members(chat_id, limit=MEMBER_CAP):
if isinstance(member.user, User):
await _save_user(ctx, member.user)
except (BadRequest, Forbidden):
return
async def _enrich_senders(client: Client, ctx: CaptureContext, chat_id: int) -> None:
rows = await ctx.pool.fetch(_MISSING_SENDERS, ctx.account_id, chat_id)
ids = [row["sender_id"] for row in rows]
for sender_id in ids:
try:
user = await client.get_users(sender_id)
except BadRequest:
continue
if isinstance(user, User):
await _save_user(ctx, user)
@register("enrich_chat")
async def enrich_chat(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
chat_id = ctx.job.params["chat_id"]
await _enrich_chat_meta(client, capture, chat_id)
await _enrich_members(client, capture, chat_id)
await _enrich_senders(client, capture, chat_id)
@@ -0,0 +1,40 @@
from io import BytesIO
from userbot.modules.avatars.repository import get_avatar_file, mark_avatar_downloaded
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
@register("fetch_avatar")
async def fetch_avatar(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
owner_kind = ctx.job.params["owner_kind"]
owner_id = ctx.job.params["owner_id"]
unique_id = ctx.job.params["unique_id"]
found = await get_avatar_file(
ctx.pool, ctx.account_id, owner_kind, owner_id, unique_id
)
if found is None:
return
file_id, downloaded = found
if downloaded or file_id is None:
return
buffer = await client.download_media(file_id, in_memory=True)
if not isinstance(buffer, BytesIO):
return
data = buffer.getvalue()
storage_key = capture.storage.put(data)
await mark_avatar_downloaded(
ctx.pool,
ctx.account_id,
owner_kind,
owner_id,
unique_id,
storage_key,
len(data),
)
@@ -0,0 +1,45 @@
from io import BytesIO
from pyrogram.types import Sticker
from userbot.modules.custom_emoji.repository import is_downloaded, upsert_downloaded
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
def _kind(sticker: Sticker) -> str:
if sticker.is_animated:
return "animated"
if sticker.is_video:
return "video"
return "static"
@register("fetch_custom_emoji")
async def fetch_custom_emoji(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
custom_emoji_id = int(ctx.job.params["custom_emoji_id"])
if await is_downloaded(ctx.pool, custom_emoji_id):
return
stickers = await client.get_custom_emoji_stickers([str(custom_emoji_id)])
if not stickers:
return
sticker = stickers[0]
buffer = await client.download_media(sticker.file_id, in_memory=True)
if not isinstance(buffer, BytesIO):
return
data = buffer.getvalue()
storage_key = capture.storage.put(data)
await upsert_downloaded(
ctx.pool,
custom_emoji_id,
storage_key,
len(data),
sticker.mime_type,
_kind(sticker),
)
@@ -0,0 +1,102 @@
import re
from pyrogram import Client, raw
from pyrogram.errors import BadRequest, Forbidden
from pyrogram.types import Chat
from userbot.modules.capture.context import CaptureContext
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
from userbot.modules.profiles.snapshots import save_chat
DEFAULT_LIMIT = 30
_USERNAME = re.compile(r"^[a-z][a-z0-9_]{3,31}$", re.IGNORECASE)
_PREFIXES = ("https://t.me/", "http://t.me/", "t.me/", "@")
def _normalize(query: str) -> str:
text = query.strip()
for prefix in _PREFIXES:
if text.lower().startswith(prefix):
text = text[len(prefix) :]
break
return text.strip("/")
_SOURCE_TYPES = (raw.types.User, raw.types.Chat, raw.types.Channel)
def _source(
peer: raw.base.Peer, users: dict, chats: dict
) -> raw.types.User | raw.types.Chat | raw.types.Channel | None:
if isinstance(peer, raw.types.PeerUser):
source = users.get(peer.user_id)
elif isinstance(peer, raw.types.PeerChannel):
source = chats.get(peer.channel_id)
elif isinstance(peer, raw.types.PeerChat):
source = chats.get(peer.chat_id)
else:
return None
return source if isinstance(source, _SOURCE_TYPES) else None
async def _save_found(
client: Client, ctx: CaptureContext, peer: raw.base.Peer, users: dict, chats: dict
) -> int | None:
source = _source(peer, users, chats)
if source is None:
return None
chat = Chat._parse_chat(client, source) # noqa: SLF001
if chat is None or chat.id is None:
return None
await save_chat(ctx, chat)
return chat.id
async def _resolve(client: Client, ctx: CaptureContext, query: str) -> int | None:
try:
chat = await client.get_chat(query)
except (BadRequest, Forbidden):
return None
if not isinstance(chat, Chat) or chat.id is None:
return None
await save_chat(ctx, chat)
return chat.id
async def _search(
client: Client, query: str, limit: int
) -> raw.base.contacts.Found | None:
try:
return await client.invoke(raw.functions.contacts.Search(q=query, limit=limit))
except (BadRequest, Forbidden):
return None
@register("search_peers")
async def search_peers(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
query = _normalize(ctx.job.params.get("query", ""))
if not query:
await ctx.report_progress({"ids": [], "done": True})
return
limit = int(ctx.job.params.get("limit", DEFAULT_LIMIT))
found = await _search(client, query, limit)
ids: list[int] = []
if found is not None:
users = {user.id: user for user in found.users}
chats = {chat.id: chat for chat in found.chats}
for peer in (*found.my_results, *found.results):
peer_id = await _save_found(client, capture, peer, users, chats)
if peer_id is not None and peer_id not in ids:
ids.append(peer_id)
if _USERNAME.match(query):
resolved = await _resolve(client, capture, query)
if resolved is not None and resolved not in ids:
ids.insert(0, resolved)
await ctx.report_progress({"ids": ids, "done": True})
@@ -0,0 +1,36 @@
from pyrogram.types import User
from userbot.modules.avatars import note_avatar
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
from userbot.modules.profiles.parse import snapshot_from_high_level
from userbot.modules.profiles.repository import write_profile
@register("sync_contacts")
async def sync_contacts(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
contacts = await client.get_contacts()
processed = 0
for user in contacts:
if not isinstance(user, User):
continue
fields, photo_file_id, photo_unique_id = snapshot_from_high_level(user)
await write_profile(ctx.pool, ctx.account_id, user.id, fields, str(user))
if photo_file_id and photo_unique_id:
await note_avatar(
ctx.pool,
ctx.account_id,
user.id,
"peer",
photo_unique_id,
photo_file_id,
)
processed += 1
await capture.contacts.refresh()
await ctx.report_progress({"processed": processed, "done": True})
@@ -0,0 +1,74 @@
from pyrogram import Client
from pyrogram.errors import BadRequest, Forbidden
from pyrogram.types import User
from userbot.modules.avatars import note_avatar
from userbot.modules.capture.context import CaptureContext
from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register
from userbot.modules.profiles.parse import snapshot_from_high_level
from userbot.modules.profiles.repository import write_profile
from userbot.modules.profiles.snapshots import save_group, save_private
SAVE_EVERY = 100
USERS_BATCH = 200
_UPSERT_DIALOG = """
INSERT INTO dialogs (account_id, chat_id) VALUES ($1, $2)
ON CONFLICT (account_id, chat_id) DO UPDATE SET updated_at = now()
"""
async def _enrich_users(client: Client, ctx: CaptureContext, ids: list[int]) -> None:
for start in range(0, len(ids), USERS_BATCH):
batch = ids[start : start + USERS_BATCH]
try:
result = await client.get_users(batch)
except (BadRequest, Forbidden):
continue
users = result if isinstance(result, list) else [result]
for user in users:
if not isinstance(user, User):
continue
fields, photo_file_id, photo_unique_id = snapshot_from_high_level(user)
await write_profile(ctx.pool, ctx.account_id, user.id, fields, str(user))
if photo_file_id and photo_unique_id:
await note_avatar(
ctx.pool,
ctx.account_id,
user.id,
"peer",
photo_unique_id,
photo_file_id,
)
@register("sync_dialogs")
async def sync_dialogs(ctx: JobContext) -> None:
client = ctx.client
if client is None:
return
capture = getattr(client, "capture", None)
if capture is None:
return
processed = ctx.job.progress.get("processed", 0)
nameless: list[int] = []
async for dialog in client.get_dialogs():
chat = dialog.chat
if chat is None or chat.id is None:
continue
chat_id = chat.id
try:
if chat_id > 0:
if not await save_private(capture, chat):
nameless.append(chat_id)
else:
await save_group(capture, chat)
except (BadRequest, Forbidden):
pass
await ctx.pool.execute(_UPSERT_DIALOG, ctx.account_id, chat_id)
processed += 1
if processed % SAVE_EVERY == 0:
await ctx.report_progress({"processed": processed})
await _enrich_users(client, capture, nameless)
await ctx.report_progress({"processed": processed, "done": True})
@@ -79,13 +79,18 @@ async def finish(
) -> None:
await pool.execute(
"UPDATE jobs SET status = $2, error = $3, finished_at = now(), "
"updated_at = now() WHERE id = $1",
"updated_at = now() WHERE id = $1 AND status = 'running'",
job_id,
status.value,
error,
)
async def is_canceled(pool: asyncpg.Pool, job_id: int) -> bool:
status = await pool.fetchval("SELECT status FROM jobs WHERE id = $1", job_id)
return status == JobStatus.CANCELED.value
async def get_job(pool: asyncpg.Pool, job_id: int) -> Job | None:
row = await pool.fetchrow("SELECT * FROM jobs WHERE id = $1", job_id)
return _row_to_job(row) if row else None
@@ -1,3 +1,7 @@
from userbot.modules.media.downloader import capture_media, self_destruct_ttl
from userbot.modules.media.downloader import (
capture_media,
media_unique_id,
self_destruct_ttl,
)
__all__ = ["capture_media", "self_destruct_ttl"]
__all__ = ["capture_media", "media_unique_id", "self_destruct_ttl"]
@@ -20,11 +20,20 @@ _MEDIA_ATTRS = (
)
_WEB_PAGE_ATTRS = ("photo", "video", "animation", "document", "audio")
def media_object(message: Message) -> tuple[str | None, Any]:
for attr in _MEDIA_ATTRS:
obj = getattr(message, attr, None)
if obj is not None:
return attr, obj
web_page = getattr(message, "web_page", None)
if web_page is not None:
for attr in _WEB_PAGE_ATTRS:
obj = getattr(web_page, attr, None)
if obj is not None:
return attr, obj
return None, None
@@ -33,6 +42,11 @@ def self_destruct_ttl(message: Message) -> int | None:
return getattr(obj, "ttl_seconds", None) if obj is not None else None
def media_unique_id(message: Message) -> str | None:
_, obj = media_object(message)
return getattr(obj, "file_unique_id", None) if obj is not None else None
async def capture_media( # noqa: PLR0913
client: Client,
message: Message,
@@ -44,6 +58,7 @@ async def capture_media( # noqa: PLR0913
kind, obj = media_object(message)
if obj is None:
return
unique_id = getattr(obj, "file_unique_id", None)
ttl = getattr(obj, "ttl_seconds", None)
want = toggles.self_destruct_media if ttl else toggles.media
file_size = getattr(obj, "file_size", None)
@@ -51,12 +66,26 @@ async def capture_media( # noqa: PLR0913
storage_key: str | None = None
downloaded = False
if want:
buffer = await client.download_media(message, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
storage_key = ctx.storage.put(data)
file_size = len(data)
existing = await repository.current_media(
ctx.pool, ctx.account_id, chat_id, message_id
)
if (
existing is not None
and existing["downloaded"]
and existing["unique_id"] == unique_id
and existing["storage_key"] is not None
):
storage_key = existing["storage_key"]
file_size = existing["file_size"]
downloaded = True
else:
target = message if getattr(message, kind or "", None) is obj else obj
buffer = await client.download_media(target, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
storage_key = ctx.storage.put(data)
file_size = len(data)
downloaded = True
await repository.insert_media(
ctx.pool,
ctx.account_id,
@@ -67,5 +96,6 @@ async def capture_media( # noqa: PLR0913
file_size,
mime,
ttl,
unique_id,
downloaded=downloaded,
)
@@ -1,7 +1,13 @@
from userbot.modules.profiles.parse import (
ProfileFields,
active_username,
snapshot_from_high_level,
snapshot_from_user,
)
__all__ = ["ProfileFields", "active_username", "snapshot_from_user"]
__all__ = [
"ProfileFields",
"active_username",
"snapshot_from_high_level",
"snapshot_from_user",
]
+33 -1
View File
@@ -1,7 +1,7 @@
from dataclasses import dataclass
from pyrogram import Client, raw
from pyrogram.types import User
from pyrogram.types import Chat, User
@dataclass(frozen=True)
@@ -37,3 +37,35 @@ def snapshot_from_user(
is_deleted_account=bool(getattr(raw_user, "deleted", False)),
)
return fields, photo_file_id, photo_unique_id
def snapshot_from_high_level(
user: User,
) -> tuple[ProfileFields, str | None, str | None]:
photo = user.photo
photo_unique_id = photo.big_photo_unique_id if photo else None
photo_file_id = photo.big_file_id if photo else None
fields = ProfileFields(
first_name=user.first_name,
last_name=user.last_name,
username=user.username,
phone=user.phone_number,
photo_unique_id=photo_unique_id,
is_deleted_account=bool(user.is_deleted),
)
return fields, photo_file_id, photo_unique_id
def snapshot_from_chat(chat: Chat) -> tuple[ProfileFields, str | None, str | None]:
photo = chat.photo
photo_unique_id = photo.big_photo_unique_id if photo else None
photo_file_id = photo.big_file_id if photo else None
fields = ProfileFields(
first_name=chat.first_name,
last_name=chat.last_name,
username=chat.username,
phone=None,
photo_unique_id=photo_unique_id,
is_deleted_account=False,
)
return fields, photo_file_id, photo_unique_id
@@ -0,0 +1,50 @@
from datetime import UTC, datetime
from pyrogram.types import Chat
from userbot.modules.avatars import note_avatar
from userbot.modules.capture.context import CaptureContext
from userbot.modules.groups.repository import insert_chat_history
from userbot.modules.profiles.parse import snapshot_from_chat
from userbot.modules.profiles.repository import write_profile
async def save_private(ctx: CaptureContext, chat: Chat) -> bool:
chat_id = chat.id or 0
fields, photo_file_id, photo_unique_id = snapshot_from_chat(chat)
await write_profile(ctx.pool, ctx.account_id, chat_id, fields, str(chat))
if photo_file_id and photo_unique_id:
await note_avatar(
ctx.pool, ctx.account_id, chat_id, "peer", photo_unique_id, photo_file_id
)
return bool(fields.first_name or fields.last_name or fields.username)
async def save_group(ctx: CaptureContext, chat: Chat) -> None:
chat_id = chat.id or 0
photo = chat.photo
photo_unique_id = photo.big_photo_unique_id if photo else None
photo_file_id = photo.big_file_id if photo else None
await insert_chat_history(
ctx.pool,
ctx.account_id,
chat_id,
0,
"meta",
chat.title,
photo_unique_id,
None,
datetime.now(UTC),
str(chat),
)
if photo_file_id and photo_unique_id:
await note_avatar(
ctx.pool, ctx.account_id, chat_id, "chat", photo_unique_id, photo_file_id
)
async def save_chat(ctx: CaptureContext, chat: Chat) -> None:
if (chat.id or 0) > 0:
await save_private(ctx, chat)
else:
await save_group(ctx, chat)
@@ -22,6 +22,20 @@ ON CONFLICT (account_id, peer_id, story_id) DO UPDATE SET
"""
async def is_downloaded(
pool: asyncpg.Pool, account_id: int, peer_id: int, story_id: int
) -> bool:
return bool(
await pool.fetchval(
"SELECT downloaded FROM stories "
"WHERE account_id = $1 AND peer_id = $2 AND story_id = $3",
account_id,
peer_id,
story_id,
)
)
async def upsert_story( # noqa: PLR0913
pool: asyncpg.Pool,
account_id: int,
@@ -0,0 +1,49 @@
from io import BytesIO
from pyrogram import Client
from pyrogram.types import Story
from userbot.modules.capture.context import CaptureContext
from userbot.modules.stories import repository
def story_peer_id(story: Story) -> int:
if story.chat is not None:
return story.chat.id or 0
if story.from_user is not None:
return story.from_user.id or 0
return 0
async def save_story(client: Client, capture: CaptureContext, story: Story) -> None:
peer_id = story_peer_id(story)
storage_key: str | None = None
file_size: int | None = None
downloaded = False
stored = await repository.is_downloaded(
capture.pool, capture.account_id, peer_id, story.id
)
if not (stored or story.deleted or story.media is None):
buffer = await client.download_media(story, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
storage_key = capture.storage.put(data)
file_size = len(data)
downloaded = True
await repository.upsert_story(
capture.pool,
capture.account_id,
peer_id,
story.id,
story.date,
story.expire_date,
story.caption,
story.media.name.lower() if story.media else None,
storage_key,
file_size,
story.views,
str(story.raw),
pinned=bool(story.pinned),
deleted=bool(story.deleted),
downloaded=downloaded,
)
+6 -2
View File
@@ -1,3 +1,7 @@
from userbot.modules.stt.service import is_transcribable, transcribe_message
from userbot.modules.stt.service import (
is_transcribable,
should_transcribe_on_backfill,
transcribe_message,
)
__all__ = ["is_transcribable", "transcribe_message"]
__all__ = ["is_transcribable", "should_transcribe_on_backfill", "transcribe_message"]
+2 -1
View File
@@ -2,7 +2,6 @@ from pyrogram import Client
from pyrogram.errors import FloodPremiumWait, FloodWait, RPCError
from userbot.modules.capture.context import CaptureContext
from userbot.modules.jobs.repository import enqueue
from userbot.modules.stt.service import transcribe_message
from utils.logging import logger
@@ -13,6 +12,8 @@ async def safe_transcribe(
try:
await transcribe_message(client, ctx, chat_id, message_id)
except (FloodWait, FloodPremiumWait):
from userbot.modules.jobs.repository import enqueue # noqa: PLC0415
await enqueue(
ctx.pool,
ctx.account_id,
@@ -9,6 +9,11 @@ UPDATE media SET extracted_text = $4
WHERE account_id = $1 AND chat_id = $2 AND message_id = $3
"""
_IS_TRANSCRIBED = """
SELECT extracted_text IS NOT NULL FROM media
WHERE account_id = $1 AND chat_id = $2 AND message_id = $3
"""
_VOICE_READS_BOX = """
SELECT md.chat_id, md.message_id, m.sender_id,
md.extracted_text IS NULL AS untranscribed
@@ -34,6 +39,12 @@ async def set_extracted_text(
await pool.execute(_SET_EXTRACTED_TEXT, account_id, chat_id, message_id, text)
async def is_transcribed(
pool: asyncpg.Pool, account_id: int, chat_id: int, message_id: int
) -> bool:
return bool(await pool.fetchval(_IS_TRANSCRIBED, account_id, chat_id, message_id))
async def voice_reads(
pool: asyncpg.Pool,
account_id: int,
+26 -1
View File
@@ -4,6 +4,7 @@ from pyrogram.types import Message
from userbot.modules.capture.context import CaptureContext
from userbot.modules.media import self_destruct_ttl
from userbot.modules.stt import repository
from utils.logging import logger
def is_transcribable(message: Message) -> bool:
@@ -12,6 +13,19 @@ def is_transcribable(message: Message) -> bool:
return message.voice is not None or message.video_note is not None
def should_transcribe_on_backfill(message: Message, self_id: int | None) -> bool:
if not is_transcribable(message):
return False
if message.outgoing:
return True
sender = message.from_user.id if message.from_user else None
if sender is None and message.sender_chat is not None:
sender = message.sender_chat.id
if sender == self_id:
return True
return not message.unread_media
async def transcribe_message(
client: Client, ctx: CaptureContext, chat_id: int, message_id: int
) -> None:
@@ -21,7 +35,18 @@ async def transcribe_message(
result = await client.invoke(
raw.functions.messages.TranscribeAudio(peer=peer, msg_id=message_id)
)
if not result.pending and result.text:
if result.pending:
logger.info(
f"[yellow]STT pending {chat_id}/{message_id} "
f"(trial_remains={result.trial_remains_num})[/]"
)
return
if result.text:
await repository.set_extracted_text(
ctx.pool, ctx.account_id, chat_id, message_id, result.text
)
else:
logger.info(
f"[yellow]STT empty {chat_id}/{message_id} "
f"(trial_remains={result.trial_remains_num})[/]"
)
+145 -91
View File
@@ -1,7 +1,7 @@
import asyncio
import contextlib
import json
from collections.abc import Callable, Coroutine
from dataclasses import dataclass
from pathlib import Path
from typing import Any
@@ -13,77 +13,156 @@ from userbot import PyroClient
from userbot.modules.capture import CaptureContext, build_capture_context
from userbot.modules.jobs import JobConsumer
from utils.env import env
from utils.jobs import enqueue
from utils.logging import logger, setup_logging
from utils.read.accounts import (
ACCOUNTS_CHANGED_CHANNEL,
inactive_session_names,
session_device_models,
sync_account,
)
from utils.read.watches import WATCHES_CHANGED_CHANNEL
from utils.storage import ContentAddressedStorage
setup_logging()
_UPSERT_ACCOUNT = """
INSERT INTO accounts
(tg_user_id, label, phone, session_name, is_active, raw, updated_at)
VALUES ($1, $2, $3, $4, TRUE, $5::jsonb, now())
ON CONFLICT (tg_user_id) DO UPDATE SET
label = EXCLUDED.label,
phone = EXCLUDED.phone,
session_name = EXCLUDED.session_name,
is_active = TRUE,
raw = EXCLUDED.raw,
updated_at = now()
RETURNING account_id
"""
@dataclass
class RunningAccount:
client: PyroClient
consumer_task: asyncio.Task
device_model: str | None
def _discover_sessions(sessions_dir: Path) -> list[Path]:
def _sessions_dir() -> Path:
sessions_dir = Path(env.tg.sessions_dir)
sessions_dir.mkdir(parents=True, exist_ok=True)
return sorted(sessions_dir.glob("*.session"))
return sessions_dir
async def _sync_account(
pool: asyncpg.Pool, client: PyroClient, session_name: str
) -> int | None:
me = client.me
if not me:
return None
raw = json.dumps(
{
"id": me.id,
"first_name": me.first_name,
"last_name": me.last_name,
"username": me.username,
"phone_number": me.phone_number,
}
async def _cancel(task: asyncio.Task) -> None:
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
async def _enqueue_once(pool: asyncpg.Pool, account_id: int, kind: str) -> None:
existing = await pool.fetchval(
"SELECT 1 FROM jobs WHERE account_id = $1 AND kind = $2 "
"AND status IN ('pending', 'running') LIMIT 1",
account_id,
kind,
)
label = " ".join(filter(None, [me.first_name, me.last_name])) or me.username
account_id = await pool.fetchval(
_UPSERT_ACCOUNT, me.id, label, me.phone_number, session_name, raw
)
logger.info(f"[green]Account synced:[/] {label} ({me.id})")
return account_id
if existing is None:
await enqueue(pool, account_id, kind, {})
logger.info(f"[green]Queued {kind}.[/]")
async def _setup_capture(
pool: asyncpg.Pool,
client: PyroClient,
account_id: int,
storage: ContentAddressedStorage,
) -> None:
client.capture = await build_capture_context(client, pool, storage, account_id)
logger.info("[green]Capture context ready.[/]")
class AccountRegistry:
def __init__(self, pool: asyncpg.Pool, storage: ContentAddressedStorage) -> None:
self._pool = pool
self._storage = storage
self._running: dict[str, RunningAccount] = {}
self._lock = asyncio.Lock()
@property
def clients(self) -> list[PyroClient]:
return [account.client for account in self._running.values()]
async def sync(self) -> None:
async with self._lock:
inactive = await inactive_session_names(self._pool)
models = await session_device_models(self._pool)
present: set[str] = set()
for path in sorted(_sessions_dir().glob("*.session")):
if path.stem in inactive:
await self._log_out(path)
continue
present.add(path.stem)
device_model = models.get(path.stem)
running = self._running.get(path.stem)
if running is not None and running.device_model != device_model:
await self._stop(path.stem)
running = None
if running is None:
await self._start(path, device_model)
for session_name in set(self._running) - present:
await self._stop(session_name)
async def close(self) -> None:
for session_name in list(self._running):
await self._stop(session_name)
async def _start(self, path: Path, device_model: str | None) -> None:
session_name = path.stem
client = PyroClient(
session_name, workdir=str(path.parent), device_model=device_model
)
try:
await client.start()
me = client.me
if me is None:
msg = f"session {session_name} is not authorized"
raise RuntimeError(msg)
account_id = await sync_account(self._pool, me, session_name)
client.capture = await build_capture_context(
client, self._pool, self._storage, account_id
)
except Exception:
logger.exception(f"[red]Failed to start session:[/] {session_name}")
with contextlib.suppress(Exception):
await client.stop()
return
consumer = JobConsumer(client, self._pool, account_id)
self._running[session_name] = RunningAccount(
client, asyncio.create_task(consumer.run()), device_model
)
logger.info(f"[green]Client started:[/] {me.full_name} ({me.id})")
await _enqueue_once(self._pool, account_id, "sync_dialogs")
await _enqueue_once(self._pool, account_id, "sync_contacts")
async def _stop(self, session_name: str) -> None:
account = self._running.pop(session_name, None)
if account is None:
return
await _cancel(account.consumer_task)
with contextlib.suppress(Exception):
await account.client.stop()
logger.info(f"[yellow]Client stopped:[/] {session_name}")
async def _log_out(self, path: Path) -> None:
account = self._running.pop(path.stem, None)
if account is not None:
await _cancel(account.consumer_task)
client = account.client
else:
client = PyroClient(
path.stem, workdir=str(path.parent), load_handlers=False
)
with contextlib.suppress(Exception):
if account is None:
await client.start()
await client.log_out()
with contextlib.suppress(Exception):
await client.stop()
path.unlink(missing_ok=True) # noqa: ASYNC240
logger.info(f"[yellow]Account logged out:[/] {path.stem}")
async def _listen_changes(
clients: list[PyroClient], tasks: set[asyncio.Task]
registry: AccountRegistry, tasks: set[asyncio.Task]
) -> asyncpg.Connection:
def spawn(coro: Coroutine[Any, Any, None]) -> None:
task = asyncio.create_task(coro)
tasks.add(task)
task.add_done_callback(tasks.discard)
def reload(
make_coro: Callable[[CaptureContext], Coroutine[Any, Any, None]],
) -> None:
for client in clients:
if client.capture is None:
continue
task = asyncio.create_task(make_coro(client.capture))
tasks.add(task)
task.add_done_callback(tasks.discard)
for client in registry.clients:
if client.capture is not None:
spawn(make_coro(client.capture))
def on_policy(
_conn: asyncpg.Connection, _pid: int, _channel: str, _payload: str
@@ -95,9 +174,15 @@ async def _listen_changes(
) -> None:
reload(lambda capture: capture.watches.refresh())
def on_accounts(
_conn: asyncpg.Connection, _pid: int, _channel: str, _payload: str
) -> None:
spawn(registry.sync())
conn = await asyncpg.connect(dsn=env.db.connection_url)
await conn.add_listener("policy_changed", on_policy)
await conn.add_listener(WATCHES_CHANGED_CHANNEL, on_watch)
await conn.add_listener(ACCOUNTS_CHANGED_CHANNEL, on_accounts)
return conn
@@ -105,52 +190,21 @@ async def runner() -> None:
pool = await container.get(asyncpg.Pool)
storage = await container.get(ContentAddressedStorage)
sessions_dir = Path(env.tg.sessions_dir)
session_files = _discover_sessions(sessions_dir)
if not session_files:
logger.warning(
f"[yellow]No .session files in {sessions_dir}/. "
f"Log in first, then restart userbot.[/]"
)
clients: list[PyroClient] = []
reload_tasks: set[asyncio.Task] = set()
consumer_tasks: list[asyncio.Task] = []
registry = AccountRegistry(pool, storage)
tasks: set[asyncio.Task] = set()
listen_conn: asyncpg.Connection | None = None
try:
for session_path in session_files:
session_name = session_path.stem
client = PyroClient(session_name, workdir=str(sessions_dir))
await client.start()
clients.append(client)
logger.info(
f"[green]Client started:[/] "
f"{client.me.full_name if client.me else 'unknown'} "
f"{client.me.id if client.me else 'unknown'}"
)
account_id = await _sync_account(pool, client, session_name)
if account_id is not None:
await _setup_capture(pool, client, account_id, storage)
consumer = JobConsumer(client, pool, account_id)
consumer_tasks.append(asyncio.create_task(consumer.run()))
if clients:
listen_conn = await _listen_changes(clients, reload_tasks)
logger.info("[green]Userbot running.[/]")
await asyncio.Event().wait()
await registry.sync()
if not registry.clients:
logger.warning("[yellow]No sessions yet. Add an account in the web UI.[/]")
listen_conn = await _listen_changes(registry, tasks)
logger.info("[green]Userbot running.[/]")
await asyncio.Event().wait()
finally:
for task in consumer_tasks:
task.cancel()
for task in consumer_tasks:
with contextlib.suppress(asyncio.CancelledError):
await task
if listen_conn is not None:
with contextlib.suppress(Exception):
await listen_conn.close()
for client in clients:
with contextlib.suppress(Exception):
await client.stop()
await registry.close()
await container.close()
+4
View File
@@ -0,0 +1,4 @@
IMMUTABLE_HEADERS = {"Cache-Control": "public, max-age=31536000, immutable"}
DAY_HEADERS = {"Cache-Control": "public, max-age=86400"}
SHORT_HEADERS = {"Cache-Control": "public, max-age=300"}
NO_STORE_HEADERS = {"Cache-Control": "no-cache"}
+66
View File
@@ -12,6 +12,7 @@ class JobStatus(StrEnum):
RUNNING = "running"
DONE = "done"
FAILED = "failed"
CANCELED = "canceled"
class Account(SQLModel, table=True):
@@ -477,3 +478,68 @@ class Alert(SQLModel, table=True):
DateTime(timezone=True), nullable=False, server_default=func.now()
)
)
class Dialog(SQLModel, table=True):
__tablename__ = "dialogs"
account_id: int = Field(primary_key=True)
chat_id: int = Field(primary_key=True)
updated_at: datetime = Field(
sa_column=Column(
DateTime(timezone=True),
nullable=False,
server_default=func.now(),
onupdate=func.now(),
)
)
class FileShare(SQLModel, table=True):
__tablename__ = "file_shares"
id: int | None = Field(default=None, primary_key=True)
account_id: int
token: str = Field(unique=True)
kind: str
storage_key: str
file_name: str
mime: str | None = None
file_size: int | None = Field(default=None, sa_column=Column(BigInteger))
title: str | None = None
chat_id: int | None = Field(default=None, sa_column=Column(BigInteger))
message_id: int | None = Field(default=None, sa_column=Column(BigInteger))
peer_id: int | None = Field(default=None, sa_column=Column(BigInteger))
story_id: int | None = Field(default=None, sa_column=Column(BigInteger))
expires_at: datetime | None = Field(
default=None, sa_column=Column(DateTime(timezone=True))
)
max_downloads: int | None = None
download_count: int = 0
last_download_at: datetime | None = Field(
default=None, sa_column=Column(DateTime(timezone=True))
)
revoked_at: datetime | None = Field(
default=None, sa_column=Column(DateTime(timezone=True))
)
created_at: datetime = Field(
sa_column=Column(
DateTime(timezone=True), nullable=False, server_default=func.now()
)
)
class FileShareHit(SQLModel, table=True):
__tablename__ = "file_share_hits"
id: int | None = Field(default=None, sa_column=Column(BigInteger, primary_key=True))
share_id: int = Field(foreign_key="file_shares.id")
ts: datetime = Field(
sa_column=Column(
DateTime(timezone=True), nullable=False, server_default=func.now()
)
)
method: str
ip: str | None = None
user_agent: str | None = None
counted: bool = False
+1
View File
@@ -34,6 +34,7 @@ class TelegramSettings(BaseSettings):
class ApiSettings(BaseSettings):
host: str = "0.0.0.0" # noqa: S104
port: int = 8080
static_dir: str = "static"
class AuthSettings(BaseSettings):
+29
View File
@@ -0,0 +1,29 @@
import json
from typing import Literal
import asyncpg
BG_EVENTS_CHANNEL = "bg_events"
EventKind = Literal["message", "edit", "delete", "reaction", "presence", "receipt"]
async def notify_bg_event( # noqa: PLR0913
pool: asyncpg.Pool,
kind: EventKind,
account_id: int,
*,
chat_id: int | None = None,
message_id: int | None = None,
message_ids: list[int] | None = None,
) -> None:
payload: dict[str, object] = {"kind": kind, "account_id": account_id}
if chat_id is not None:
payload["chat_id"] = chat_id
if message_id is not None:
payload["message_id"] = message_id
if message_ids is not None:
payload["message_ids"] = message_ids
await pool.execute(
"SELECT pg_notify($1, $2)", BG_EVENTS_CHANNEL, json.dumps(payload)
)
+231
View File
@@ -0,0 +1,231 @@
import re
from datetime import UTC, datetime, timedelta
from urllib.parse import quote
_MIME_EXTENSIONS = {
"application/pdf": ".pdf",
"application/x-tgsticker": ".tgs",
"application/zip": ".zip",
"audio/mpeg": ".mp3",
"audio/ogg": ".ogg",
"image/gif": ".gif",
"image/jpeg": ".jpg",
"image/png": ".png",
"image/webp": ".webp",
"text/plain": ".txt",
"video/mp4": ".mp4",
"video/quicktime": ".mov",
"video/webm": ".webm",
}
_KIND_EXTENSIONS = {
"animation": ".mp4",
"audio": ".mp3",
"gif": ".mp4",
"photo": ".jpg",
"sticker": ".webp",
"video": ".mp4",
"video_note": ".mp4",
"voice": ".ogg",
}
_EXTENSION_MIMES = {
".flac": "audio/flac",
".gif": "image/gif",
".heic": "image/heic",
".jpeg": "image/jpeg",
".jpg": "image/jpeg",
".m4a": "audio/mp4",
".mov": "video/quicktime",
".mp3": "audio/mpeg",
".mp4": "video/mp4",
".ogg": "audio/ogg",
".pdf": "application/pdf",
".png": "image/png",
".svg": "image/svg+xml",
".txt": "text/plain",
".wav": "audio/wav",
".webm": "video/webm",
".webp": "image/webp",
}
_KIND_MIMES = {
"animation": "video/mp4",
"gif": "video/mp4",
"photo": "image/jpeg",
"video": "video/mp4",
"video_note": "video/mp4",
"voice": "audio/ogg",
}
_GENERIC_MIMES = {
"application/octet-stream",
"application/binary",
"binary/octet-stream",
"",
}
_INLINE_MIME_PREFIXES = ("image/", "video/", "audio/", "text/")
_INLINE_MIMES = {"application/pdf", "application/json"}
_PREVIEW_AGENTS = (
"telegrambot",
"twitterbot",
"facebookexternalhit",
"whatsapp",
"discordbot",
"slackbot",
"skypeuripreview",
"vkshare",
"redditbot",
"linkedinbot",
"embedly",
"quora link preview",
"googlebot",
"bingbot",
"applebot",
"yandexbot",
"duckduckbot",
"petalbot",
"ahrefsbot",
"semrushbot",
"headlesschrome",
"python-requests",
)
_TRANSLIT = {
"а": "a",
"б": "b",
"в": "v",
"г": "g",
"д": "d",
"е": "e",
"ё": "e",
"ж": "zh",
"з": "z",
"и": "i",
"й": "y",
"к": "k",
"л": "l",
"м": "m",
"н": "n",
"о": "o",
"п": "p",
"р": "r",
"с": "s",
"т": "t",
"у": "u",
"ф": "f",
"х": "h",
"ц": "c",
"ч": "ch",
"ш": "sh",
"щ": "sch",
"ъ": "",
"ы": "y",
"ь": "",
"э": "e",
"ю": "yu",
"я": "ya",
}
_URL_UNSAFE = re.compile(r"[^A-Za-z0-9._-]+")
_URL_REPEATS = re.compile(r"_{2,}")
_UNSAFE_CHARS = re.compile(r'[\\/:*?"<>|\x00-\x1f]+')
_SPACES = re.compile(r"\s+")
_NAME_LIMIT = 120
def extension_for(kind: str | None, mime: str | None) -> str:
if mime and mime in _MIME_EXTENSIONS:
return _MIME_EXTENSIONS[mime]
if kind and kind in _KIND_EXTENSIONS:
return _KIND_EXTENSIONS[kind]
if mime and "/" in mime:
tail = mime.rsplit("/", 1)[1].split(";")[0].strip()
if tail.isalnum():
return f".{tail}"
return ".bin"
def sanitize_name(name: str) -> str:
cleaned = _SPACES.sub(" ", _UNSAFE_CHARS.sub("_", name)).strip(" .")
if len(cleaned) > _NAME_LIMIT:
head, dot, tail = cleaned.rpartition(".")
cleaned = (
f"{head[: _NAME_LIMIT - len(tail) - 1]}{dot}{tail}"
if dot and len(tail) < 12 # noqa: PLR2004
else cleaned[:_NAME_LIMIT]
)
return cleaned or "file"
def url_slug(file_name: str) -> str:
lowered = "".join(
_TRANSLIT.get(ch, _TRANSLIT.get(ch.lower(), ch)) for ch in file_name
)
slug = _URL_REPEATS.sub("_", _URL_UNSAFE.sub("_", lowered)).strip("_.")
if not slug:
return "file"
return slug[:_NAME_LIMIT]
def media_file_name(
kind: str | None, mime: str | None, message_id: int, original: str | None = None
) -> str:
if original:
name = sanitize_name(original)
return name if "." in name else f"{name}{extension_for(kind, mime)}"
return f"{kind or 'media'}_{message_id}{extension_for(kind, mime)}"
def story_file_name(peer_id: int, story_id: int, media_kind: str | None) -> str:
return f"story_{peer_id}_{story_id}{extension_for(media_kind or 'photo', None)}"
def resolve_mime(
kind: str | None, mime: str | None, file_name: str | None = None
) -> str:
if mime and mime.lower() not in _GENERIC_MIMES:
return mime
if file_name and "." in file_name:
by_extension = _EXTENSION_MIMES.get(f".{file_name.rsplit('.', 1)[1].lower()}")
if by_extension:
return by_extension
return _KIND_MIMES.get(kind or "", mime or "application/octet-stream")
def is_inline_mime(mime: str | None) -> bool:
if not mime:
return False
return mime in _INLINE_MIMES or mime.startswith(_INLINE_MIME_PREFIXES)
def content_disposition(file_name: str, *, attachment: bool) -> str:
kind = "attachment" if attachment else "inline"
ascii_name = file_name.encode("ascii", "replace").decode("ascii").replace('"', "_")
return f"{kind}; filename=\"{ascii_name}\"; filename*=UTF-8''{quote(file_name)}"
def is_preview_agent(user_agent: str | None) -> bool:
if not user_agent:
return False
lowered = user_agent.lower()
return any(marker in lowered for marker in _PREVIEW_AGENTS)
def counts_as_download(
method: str, user_agent: str | None, range_header: str | None
) -> bool:
if method.upper() != "GET" or is_preview_agent(user_agent):
return False
if range_header is None:
return True
return range_header.replace(" ", "").startswith("bytes=0-")
def expiry_from_seconds(seconds: int | None) -> datetime | None:
if seconds is None or seconds <= 0:
return None
return datetime.now(UTC) + timedelta(seconds=seconds)
+19
View File
@@ -16,3 +16,22 @@ DEFAULTS: dict[ChatKind, CaptureToggles] = {
backfill=True,
),
}
TRACKING: dict[ChatKind, CaptureToggles] = {
ChatKind.CHANNEL: CaptureToggles(
messages=True,
media=True,
reactions=True,
track_edits_deletes=True,
backfill=True,
),
ChatKind.GROUP: CaptureToggles(
messages=True,
media=True,
reactions=True,
track_edits_deletes=True,
profile_history=True,
backfill=True,
),
ChatKind.DM: DEFAULTS[ChatKind.DM],
}
+46 -3
View File
@@ -85,6 +85,18 @@ async def list_folders(pool: asyncpg.Pool, account_id: int) -> list[FolderSpec]:
return [_row_to_folder(row) for row in rows]
async def get_folder(
pool: asyncpg.Pool, account_id: int, folder_id: int
) -> FolderSpec | None:
row = await pool.fetchrow(
"SELECT folder_id, title, order_index, is_chatlist, raw "
"FROM folders WHERE account_id = $1 AND folder_id = $2",
account_id,
folder_id,
)
return _row_to_folder(row) if row else None
async def create_policy(
pool: asyncpg.Pool,
account_id: int | None,
@@ -110,10 +122,24 @@ async def get_policy(pool: asyncpg.Pool, policy_id: int) -> PolicyRecord | None:
return PolicyRecord(**dict(row)) if row else None
async def find_policy(
pool: asyncpg.Pool, account_id: int, scope_type: ScopeType, scope_id: int | None
) -> PolicyRecord | None:
row = await pool.fetchrow(
"SELECT * FROM capture_policy WHERE account_id = $1 AND scope_type = $2 "
"AND scope_id IS NOT DISTINCT FROM $3",
account_id,
scope_type.value,
scope_id,
)
return PolicyRecord(**dict(row)) if row else None
async def list_policies(pool: asyncpg.Pool, account_id: int) -> list[PolicyRecord]:
rows = await pool.fetch(
"SELECT * FROM capture_policy WHERE account_id = $1 OR account_id IS NULL "
"ORDER BY scope_type, scope_id",
"SELECT DISTINCT ON (scope_type, scope_id) * FROM capture_policy "
"WHERE account_id = $1 OR account_id IS NULL "
"ORDER BY scope_type, scope_id, account_id NULLS LAST",
account_id,
)
return [PolicyRecord(**dict(row)) for row in rows]
@@ -131,6 +157,22 @@ async def update_policy(
return PolicyRecord(**dict(row)) if row else None
async def override_policy(
pool: asyncpg.Pool, policy_id: int, account_id: int, toggles: CaptureToggles
) -> PolicyRecord | None:
record = await get_policy(pool, policy_id)
if record is None:
return None
if record.account_id == account_id:
return await update_policy(pool, policy_id, toggles)
existing = await find_policy(pool, account_id, record.scope_type, record.scope_id)
if existing is not None:
return await update_policy(pool, existing.id, toggles)
return await create_policy(
pool, account_id, record.scope_type, record.scope_id, toggles
)
async def delete_policy(pool: asyncpg.Pool, policy_id: int) -> bool:
result = await pool.execute("DELETE FROM capture_policy WHERE id = $1", policy_id)
return result.endswith("1")
@@ -138,7 +180,8 @@ async def delete_policy(pool: asyncpg.Pool, policy_id: int) -> bool:
async def load_policy_set(pool: asyncpg.Pool, account_id: int) -> PolicySet:
rows = await pool.fetch(
"SELECT * FROM capture_policy WHERE account_id = $1 OR account_id IS NULL",
"SELECT * FROM capture_policy WHERE account_id = $1 OR account_id IS NULL "
"ORDER BY account_id NULLS FIRST",
account_id,
)
policies = PolicySet()
+98
View File
@@ -0,0 +1,98 @@
import json
import asyncpg
from pyrogram.types import User
from utils.read.models import AccountView
ACCOUNTS_CHANGED_CHANNEL = "accounts_changed"
_ACCOUNT_COLS = "account_id, label, phone, tg_user_id, is_active, device_model"
_UPSERT_ACCOUNT = """
INSERT INTO accounts
(tg_user_id, label, phone, session_name, is_active, raw, updated_at)
VALUES ($1, $2, $3, $4, TRUE, $5::jsonb, now())
ON CONFLICT (tg_user_id) DO UPDATE SET
label = EXCLUDED.label,
phone = EXCLUDED.phone,
session_name = EXCLUDED.session_name,
is_active = TRUE,
raw = EXCLUDED.raw,
updated_at = now()
RETURNING account_id
"""
_SET_DEVICE_MODEL = f"""
UPDATE accounts SET device_model = $2, updated_at = now()
WHERE account_id = $1
RETURNING {_ACCOUNT_COLS}
""" # noqa: S608
async def self_user_id(pool: asyncpg.Pool, account_id: int) -> int | None:
return await pool.fetchval(
"SELECT tg_user_id FROM accounts WHERE account_id = $1", account_id
)
async def list_accounts(pool: asyncpg.Pool) -> list[AccountView]:
rows = await pool.fetch(
f"SELECT {_ACCOUNT_COLS} FROM accounts ORDER BY account_id" # noqa: S608
)
return [AccountView(**dict(row)) for row in rows]
async def get_account(pool: asyncpg.Pool, account_id: int) -> AccountView | None:
row = await pool.fetchrow(
f"SELECT {_ACCOUNT_COLS} FROM accounts WHERE account_id = $1", # noqa: S608
account_id,
)
return AccountView(**dict(row)) if row else None
async def sync_account(pool: asyncpg.Pool, me: User, session_name: str) -> int:
raw = json.dumps(
{
"id": me.id,
"first_name": me.first_name,
"last_name": me.last_name,
"username": me.username,
"phone_number": me.phone_number,
}
)
label = " ".join(filter(None, [me.first_name, me.last_name])) or me.username
return await pool.fetchval(
_UPSERT_ACCOUNT, me.id, label, me.phone_number, session_name, raw
)
async def set_device_model(
pool: asyncpg.Pool, account_id: int, device_model: str
) -> AccountView | None:
row = await pool.fetchrow(_SET_DEVICE_MODEL, account_id, device_model)
return AccountView(**dict(row)) if row else None
async def session_device_models(pool: asyncpg.Pool) -> dict[str, str]:
rows = await pool.fetch(
"SELECT session_name, device_model FROM accounts WHERE device_model IS NOT NULL"
)
return {row["session_name"]: row["device_model"] for row in rows}
async def deactivate_account(pool: asyncpg.Pool, account_id: int) -> str | None:
return await pool.fetchval(
"UPDATE accounts SET is_active = FALSE, updated_at = now() "
"WHERE account_id = $1 RETURNING session_name",
account_id,
)
async def inactive_session_names(pool: asyncpg.Pool) -> set[str]:
rows = await pool.fetch("SELECT session_name FROM accounts WHERE NOT is_active")
return {row["session_name"] for row in rows}
async def notify_accounts_changed(pool: asyncpg.Pool) -> None:
await pool.execute("SELECT pg_notify($1, '')", ACCOUNTS_CHANGED_CHANNEL)
+69
View File
@@ -0,0 +1,69 @@
from datetime import UTC, datetime, timedelta
import asyncpg
from utils.read.accounts import self_user_id
from utils.read.models import ResponseStats, VolumeBucket
async def message_volume(
pool: asyncpg.Pool, account_id: int, chat_id: int, *, days: int = 90
) -> list[VolumeBucket]:
self_id = await self_user_id(pool, account_id)
rows = await pool.fetch(
"SELECT date_trunc('day', date) AS bucket, count(*) AS total, "
"count(*) FILTER (WHERE sender_id = $3) AS outgoing "
"FROM messages "
"WHERE account_id = $1 AND chat_id = $2 AND date >= $4 "
"GROUP BY bucket ORDER BY bucket",
account_id,
chat_id,
self_id,
datetime.now(UTC) - timedelta(days=days),
)
return [
VolumeBucket(
bucket=row["bucket"],
total=row["total"],
outgoing=row["outgoing"],
incoming=row["total"] - row["outgoing"],
)
for row in rows
]
async def response_stats(
pool: asyncpg.Pool, account_id: int, chat_id: int
) -> ResponseStats:
self_id = await self_user_id(pool, account_id)
rows = await pool.fetch(
"WITH ordered AS ("
"SELECT sender_id, date, "
"lag(sender_id) OVER w AS prev_sender, "
"lag(date) OVER w AS prev_date "
"FROM messages "
"WHERE account_id = $1 AND chat_id = $2 AND sender_id IS NOT NULL "
"WINDOW w AS (ORDER BY date, message_id)), "
"resp AS ("
"SELECT (sender_id = $3) AS is_mine, "
"EXTRACT(EPOCH FROM (date - prev_date)) AS secs "
"FROM ordered "
"WHERE prev_sender IS NOT NULL AND prev_sender <> sender_id) "
"SELECT is_mine, count(*) AS n, "
"percentile_cont(0.5) WITHIN GROUP (ORDER BY secs) AS median_secs "
"FROM resp GROUP BY is_mine",
account_id,
chat_id,
self_id,
)
stats = ResponseStats(
mine_median_seconds=None, mine_count=0, their_median_seconds=None, their_count=0
)
for row in rows:
if row["is_mine"]:
stats.mine_median_seconds = row["median_secs"]
stats.mine_count = row["n"]
else:
stats.their_median_seconds = row["median_secs"]
stats.their_count = row["n"]
return stats
+50
View File
@@ -0,0 +1,50 @@
import asyncpg
from utils.read.models import AvatarHistoryView, AvatarRef
_PEER_UNIQUE_ID = """
SELECT photo_unique_id FROM peers
WHERE account_id = $1 AND peer_id = $2
"""
_CHAT_UNIQUE_ID = """
SELECT photo_unique_id FROM chat_history
WHERE account_id = $1 AND chat_id = $2 AND photo_unique_id IS NOT NULL
ORDER BY ts DESC LIMIT 1
"""
_AVATAR = """
SELECT unique_id, storage_key, downloaded, mime FROM avatars
WHERE account_id = $1 AND owner_id = $2 AND unique_id = $3
"""
_AVATAR_HISTORY = """
SELECT unique_id, first_seen_at, downloaded FROM avatars
WHERE account_id = $1 AND owner_id = $2
ORDER BY first_seen_at DESC
"""
async def current_avatar(
pool: asyncpg.Pool, account_id: int, owner_kind: str, owner_id: int
) -> AvatarRef | None:
query = _PEER_UNIQUE_ID if owner_kind == "peer" else _CHAT_UNIQUE_ID
unique_id = await pool.fetchval(query, account_id, owner_id)
if unique_id is None:
return None
row = await pool.fetchrow(_AVATAR, account_id, owner_id, unique_id)
return AvatarRef(**dict(row)) if row else None
async def avatar_by_unique_id(
pool: asyncpg.Pool, account_id: int, owner_id: int, unique_id: str
) -> AvatarRef | None:
row = await pool.fetchrow(_AVATAR, account_id, owner_id, unique_id)
return AvatarRef(**dict(row)) if row else None
async def avatar_history(
pool: asyncpg.Pool, account_id: int, owner_id: int
) -> list[AvatarHistoryView]:
rows = await pool.fetch(_AVATAR_HISTORY, account_id, owner_id)
return [AvatarHistoryView(**dict(row)) for row in rows]
+257 -46
View File
@@ -1,76 +1,280 @@
import asyncpg
from utils.read.models import ChatListItem, MessageVersionView, MessageView, Page
from utils.policy.models import FolderSpec
from utils.read.accounts import self_user_id
from utils.read.message_view import build_message_view, load_raw, media_ref_from
from utils.read.models import (
ChatListItem,
MediaRef,
MessageVersionView,
MessageView,
Page,
)
from utils.read.read_receipts import read_up_to
_MESSAGE_COLS = (
"chat_id, message_id, date, sender_id, text, "
"has_media, is_self_destruct, edited_at, deleted_at"
"chat_id, message_id, date, sender_id, text, has_media, is_self_destruct, "
"edited_at, deleted_at, raw, raw->>'media_group_id' AS media_group_id"
)
def _peer_title(
first: str | None, last: str | None, username: str | None
) -> str | None:
name = " ".join(part for part in (first, last) if part)
return name or username
async def _media_map(
pool: asyncpg.Pool, account_id: int, rows: list[asyncpg.Record]
) -> dict[tuple[int, int], asyncpg.Record]:
message_ids = list({row["message_id"] for row in rows})
if not message_ids:
return {}
media_rows = await pool.fetch(
"SELECT id, chat_id, message_id, kind, downloaded, mime, file_size, "
"ttl_seconds, extracted_text FROM media "
"WHERE account_id = $1 AND message_id = ANY($2::bigint[])",
account_id,
message_ids,
)
return {(row["chat_id"], row["message_id"]): row for row in media_rows}
def _single_media(
row: asyncpg.Record, raw: dict, media_by_key: dict[tuple[int, int], asyncpg.Record]
) -> list[MediaRef]:
media_row = media_by_key.get((row["chat_id"], row["message_id"]))
if not (row["has_media"] or media_row):
return []
ref = media_ref_from(row["message_id"], raw, media_row)
return [ref] if ref else []
_ALL_IDS = """
SELECT chat_id FROM chat_stats WHERE account_id = $1
UNION
SELECT chat_id FROM dialogs WHERE account_id = $1
UNION
SELECT scope_id AS chat_id FROM capture_policy
WHERE account_id = $1 AND scope_type = 'chat' AND scope_id IS NOT NULL
"""
_ONE_ID = """
SELECT chat_id FROM chat_stats WHERE account_id = $1 AND chat_id = $2
UNION
SELECT chat_id FROM dialogs WHERE account_id = $1 AND chat_id = $2
UNION
SELECT scope_id AS chat_id FROM capture_policy
WHERE account_id = $1 AND scope_type = 'chat' AND scope_id = $2
"""
_CHAT_ROWS = """
WITH ids AS ({ids})
SELECT ids.chat_id,
COALESCE(cs.message_count, 0) AS message_count,
cs.last_date, cs.last_text, cs.last_sender_id,
COALESCE(named.title,
NULLIF(trim(concat_ws(' ', p.first_name, p.last_name)), ''),
p.username) AS title,
COALESCE(typed.is_broadcast, false) AS is_broadcast,
COALESCE((p.raw->>'is_bot')::bool, (p.raw->>'bot')::bool,
p.raw->>'type' = 'ChatType.BOT', false) AS is_bot,
COALESCE((p.raw->>'is_contact')::bool, (p.raw->>'contact')::bool,
false) AS is_contact,
EXISTS (SELECT 1 FROM avatars a
WHERE a.account_id = $1 AND a.owner_id = ids.chat_id) AS has_avatar
FROM ids
LEFT JOIN chat_stats cs ON cs.account_id = $1 AND cs.chat_id = ids.chat_id
LEFT JOIN peers p ON p.account_id = $1 AND p.peer_id = ids.chat_id
LEFT JOIN LATERAL (
SELECT ch.title FROM chat_history ch
WHERE ch.account_id = $1 AND ch.chat_id = ids.chat_id AND ch.title IS NOT NULL
ORDER BY ch.ts DESC LIMIT 1
) named ON true
LEFT JOIN LATERAL (
SELECT COALESCE(ch.raw->'chat'->>'type', ch.raw->>'type')
= 'ChatType.CHANNEL' AS is_broadcast
FROM chat_history ch
WHERE ch.account_id = $1 AND ch.chat_id = ids.chat_id
AND COALESCE(ch.raw->'chat'->>'type', ch.raw->>'type') IS NOT NULL
ORDER BY ch.ts DESC LIMIT 1
) typed ON true
"""
def _folder_filter(base: int) -> str:
return (
f"NOT (chat.chat_id = ANY(${base + 1}::bigint[])) "
f"AND (chat.chat_id = ANY(${base + 2}::bigint[]) "
f"OR (NOT ${base + 3}::bool AND CASE "
f"WHEN chat.is_broadcast THEN ${base + 4}::bool "
f"WHEN chat.chat_id < 0 THEN ${base + 5}::bool "
f"WHEN chat.is_bot THEN ${base + 6}::bool "
f"WHEN chat.is_contact THEN ${base + 7}::bool "
f"ELSE ${base + 8}::bool END))"
)
def _folder_params(folder: FolderSpec) -> list[object]:
return [
sorted(folder.exclude_ids),
sorted(folder.include_ids | folder.pinned_ids),
folder.is_chatlist,
folder.broadcasts,
folder.groups,
folder.bots,
folder.contacts,
folder.non_contacts,
]
def _chat_item(row: asyncpg.Record) -> ChatListItem:
return ChatListItem(
chat_id=row["chat_id"],
title=row["title"],
kind="private" if row["chat_id"] > 0 else "group",
has_avatar=row["has_avatar"],
is_bot=bool(row["is_bot"]),
is_contact=bool(row["is_contact"]),
is_broadcast=bool(row["is_broadcast"]),
message_count=row["message_count"],
last_date=row["last_date"],
last_text=row["last_text"],
last_sender_id=row["last_sender_id"],
)
async def list_chats(
pool: asyncpg.Pool, account_id: int, page: Page
pool: asyncpg.Pool,
account_id: int,
page: Page,
*,
folder: FolderSpec | None = None,
search: str | None = None,
) -> list[ChatListItem]:
rows = await pool.fetch(
"SELECT m.chat_id, count(*) AS message_count, max(m.date) AS last_date, "
"(SELECT p.first_name FROM peers p "
"WHERE p.account_id = $1 AND p.peer_id = m.chat_id) AS first_name, "
"(SELECT p.last_name FROM peers p "
"WHERE p.account_id = $1 AND p.peer_id = m.chat_id) AS last_name, "
"(SELECT p.username FROM peers p "
"WHERE p.account_id = $1 AND p.peer_id = m.chat_id) AS username, "
"(SELECT ch.title FROM chat_history ch "
"WHERE ch.account_id = $1 AND ch.chat_id = m.chat_id "
"AND ch.title IS NOT NULL ORDER BY ch.ts DESC LIMIT 1) AS group_title "
"FROM messages m WHERE m.account_id = $1 "
"GROUP BY m.chat_id ORDER BY last_date DESC LIMIT $2 OFFSET $3",
account_id,
page.capped_limit,
page.offset,
params: list[object] = [account_id, page.capped_limit, page.offset]
rows_sql = _CHAT_ROWS.format(ids=_ALL_IDS)
clauses: list[str] = []
if folder is not None:
clauses.append(_folder_filter(len(params)))
params.extend(_folder_params(folder))
if search:
params.append(f"%{search}%")
clauses.append(f"chat.title ILIKE ${len(params)}")
where = f" WHERE {' AND '.join(clauses)}" if clauses else ""
query = (
f"SELECT chat.* FROM ({rows_sql}) chat{where} " # noqa: S608
"ORDER BY last_date DESC NULLS LAST, chat_id DESC LIMIT $2 OFFSET $3"
)
items = []
for row in rows:
title = row["group_title"] or _peer_title(
row["first_name"], row["last_name"], row["username"]
)
items.append(
ChatListItem(
chat_id=row["chat_id"],
title=title,
message_count=row["message_count"],
last_date=row["last_date"],
)
)
return items
rows = await pool.fetch(query, *params)
return [_chat_item(row) for row in rows]
async def get_chat_history(
async def get_chat(
pool: asyncpg.Pool, account_id: int, chat_id: int
) -> ChatListItem | None:
row = await pool.fetchrow(_CHAT_ROWS.format(ids=_ONE_ID), account_id, chat_id)
return _chat_item(row) if row is not None else None
async def get_chat_history( # noqa: PLR0913
pool: asyncpg.Pool,
account_id: int,
chat_id: int,
page: Page,
*,
include_deleted: bool = True,
before_id: int | None = None,
after_id: int | None = None,
) -> list[MessageView]:
where = "account_id = $1 AND chat_id = $2"
if not include_deleted:
where += " AND deleted_at IS NULL"
rows = await pool.fetch(
params: list[object] = [account_id, chat_id]
if after_id is not None:
params.append(after_id)
where += f" AND message_id > ${len(params)}"
order = "date ASC, message_id ASC"
elif before_id is not None:
params.append(before_id)
where += f" AND message_id < ${len(params)}"
order = "date DESC, message_id DESC"
else:
order = "date DESC, message_id DESC"
params.append(page.capped_limit)
query = (
f"SELECT {_MESSAGE_COLS} FROM messages WHERE {where} " # noqa: S608
"ORDER BY date DESC, message_id DESC LIMIT $3 OFFSET $4",
f"ORDER BY {order} LIMIT ${len(params)}"
)
if before_id is None and after_id is None:
params.append(page.offset)
query += f" OFFSET ${len(params)}"
rows = await pool.fetch(query, *params)
media_by_key = await _media_map(pool, account_id, rows)
parsed = [(row, load_raw(row["raw"])) for row in rows]
views: list[MessageView] = []
index = 0
while index < len(parsed):
group_id = parsed[index][0]["media_group_id"]
end = index + 1
if group_id is not None:
while end < len(parsed) and parsed[end][0]["media_group_id"] == group_id:
end += 1
members = parsed[index:end]
if len(members) == 1:
row, raw = members[0]
views.append(
build_message_view(row, raw, _single_media(row, raw, media_by_key))
)
else:
views.append(_build_album(members, media_by_key))
index = end
await _apply_read_status(pool, account_id, chat_id, views)
return views
async def get_message(
pool: asyncpg.Pool, account_id: int, chat_id: int, message_id: int
) -> MessageView | None:
row = await pool.fetchrow(
f"SELECT {_MESSAGE_COLS} FROM messages " # noqa: S608
"WHERE account_id = $1 AND chat_id = $2 AND message_id = $3",
account_id,
chat_id,
page.capped_limit,
page.offset,
message_id,
)
return [MessageView(**dict(row)) for row in rows]
if row is None:
return None
media_by_key = await _media_map(pool, account_id, [row])
raw = load_raw(row["raw"])
view = build_message_view(row, raw, _single_media(row, raw, media_by_key))
await _apply_read_status(pool, account_id, chat_id, [view])
return view
async def _apply_read_status(
pool: asyncpg.Pool, account_id: int, chat_id: int, views: list[MessageView]
) -> None:
self_id = await self_user_id(pool, account_id)
if self_id is None:
return
marker = await read_up_to(pool, account_id, chat_id)
if marker is None:
return
for view in views:
if view.sender_id == self_id and view.message_id <= marker:
view.read = True
def _build_album(
members: list[tuple[asyncpg.Record, dict]],
media_by_key: dict[tuple[int, int], asyncpg.Record],
) -> MessageView:
ordered = sorted(members, key=lambda m: m[0]["message_id"])
media: list[MediaRef] = []
for row, raw in ordered:
media_row = media_by_key.get((row["chat_id"], row["message_id"]))
ref = media_ref_from(row["message_id"], raw, media_row)
if ref:
media.append(ref)
primary_row, primary_raw = next(
((row, raw) for row, raw in ordered if row["text"]), ordered[0]
)
return build_message_view(primary_row, primary_raw, media)
async def get_deleted_messages(
@@ -88,7 +292,14 @@ async def get_deleted_messages(
f"ORDER BY deleted_at DESC LIMIT ${len(params) - 1} OFFSET ${len(params)}",
*params,
)
return [MessageView(**dict(row)) for row in rows]
media_by_key = await _media_map(pool, account_id, rows)
views: list[MessageView] = []
for row in rows:
raw = load_raw(row["raw"])
views.append(
build_message_view(row, raw, _single_media(row, raw, media_by_key))
)
return views
async def get_message_versions(
+15
View File
@@ -0,0 +1,15 @@
import asyncpg
from utils.read.models import CustomEmojiRef
_GET = """
SELECT storage_key, downloaded, mime, kind FROM custom_emoji
WHERE custom_emoji_id = $1
"""
async def current_custom_emoji(
pool: asyncpg.Pool, custom_emoji_id: int
) -> CustomEmojiRef | None:
row = await pool.fetchrow(_GET, custom_emoji_id)
return CustomEmojiRef(**dict(row)) if row else None
+141
View File
@@ -0,0 +1,141 @@
import asyncpg
from utils.policy.models import ChatKind
from utils.read.models import DiscoverItem
_ESCAPE = str.maketrans({"\\": "\\\\", "%": r"\%", "_": r"\_"})
_IS_BROADCAST = """
SELECT COALESCE(raw->'chat'->>'type', raw->>'type') = 'ChatType.CHANNEL'
FROM chat_history
WHERE account_id = $1 AND chat_id = $2
AND COALESCE(raw->'chat'->>'type', raw->>'type') IS NOT NULL
ORDER BY ts DESC LIMIT 1
"""
_IS_TRACKED = """
SELECT EXISTS (
SELECT 1 FROM capture_policy
WHERE account_id = $1 AND scope_type = 'chat' AND scope_id = $2
)
"""
_ITEMS = """
WITH chat_meta AS (
SELECT DISTINCT ON (chat_id) chat_id, title,
COALESCE(raw->'chat'->>'username', raw->>'username') AS username,
COALESCE(raw->'chat'->>'type', raw->>'type') = 'ChatType.CHANNEL'
AS is_broadcast
FROM chat_history
WHERE account_id = $1 AND title IS NOT NULL
ORDER BY chat_id, ts DESC
), hits AS (
SELECT p.peer_id AS chat_id,
COALESCE(NULLIF(concat_ws(' ', p.first_name, p.last_name), ''), p.username)
AS title,
p.username,
COALESCE((p.raw->>'is_bot')::bool, (p.raw->>'bot')::bool, false) AS is_bot,
COALESCE((p.raw->>'is_contact')::bool, (p.raw->>'contact')::bool, false)
AS is_contact,
false AS is_broadcast
FROM peers p
WHERE p.account_id = $1 AND (p.peer_id = ANY($3::bigint[]) OR ($2 <> '' AND (
concat_ws(' ', p.first_name, p.last_name) ILIKE $2
OR p.username ILIKE $2 OR p.phone ILIKE $2)))
UNION ALL
SELECT c.chat_id, c.title, c.username, false, false,
COALESCE(c.is_broadcast, false)
FROM chat_meta c
WHERE c.chat_id = ANY($3::bigint[])
OR ($2 <> '' AND (c.title ILIKE $2 OR c.username ILIKE $2))
), merged AS (
SELECT chat_id, max(title) AS title, max(username) AS username,
bool_or(is_bot) AS is_bot, bool_or(is_contact) AS is_contact,
bool_or(is_broadcast) AS is_broadcast
FROM hits GROUP BY chat_id
), counts AS (
SELECT chat_id, count(*) AS message_count FROM messages
WHERE account_id = $1 AND chat_id IN (SELECT chat_id FROM merged)
GROUP BY chat_id
)
SELECT m.chat_id, m.title, m.username, m.is_bot, m.is_contact, m.is_broadcast,
COALESCE(c.message_count, 0) AS message_count,
EXISTS (SELECT 1 FROM avatars a
WHERE a.account_id = $1 AND a.owner_id = m.chat_id) AS has_avatar,
EXISTS (SELECT 1 FROM dialogs d
WHERE d.account_id = $1 AND d.chat_id = m.chat_id) AS in_dialogs,
EXISTS (SELECT 1 FROM capture_policy cp WHERE cp.account_id = $1
AND cp.scope_type = 'chat' AND cp.scope_id = m.chat_id) AS tracked
FROM merged m LEFT JOIN counts c ON c.chat_id = m.chat_id
ORDER BY in_dialogs DESC, message_count DESC, is_contact DESC, m.title
LIMIT $4
"""
def _kind(chat_id: int, *, is_broadcast: bool) -> str:
if chat_id > 0:
return "private"
return "channel" if is_broadcast else "group"
def _to_item(row: asyncpg.Record) -> DiscoverItem:
return DiscoverItem(
chat_id=row["chat_id"],
title=row["title"],
username=row["username"],
kind=_kind(row["chat_id"], is_broadcast=row["is_broadcast"]),
is_bot=row["is_bot"],
is_contact=row["is_contact"],
has_avatar=row["has_avatar"],
message_count=row["message_count"],
in_dialogs=row["in_dialogs"],
tracked=row["tracked"],
)
async def search(
pool: asyncpg.Pool, account_id: int, query: str, limit: int
) -> list[DiscoverItem]:
text = query.strip()
if not text:
return []
rows = await pool.fetch(
_ITEMS, account_id, f"%{text.translate(_ESCAPE)}%", [], limit
)
return [_to_item(row) for row in rows]
async def by_ids(
pool: asyncpg.Pool, account_id: int, ids: list[int]
) -> list[DiscoverItem]:
if not ids:
return []
rows = await pool.fetch(_ITEMS, account_id, "", ids, len(ids))
by_id = {row["chat_id"]: _to_item(row) for row in rows}
return [by_id[chat_id] for chat_id in ids if chat_id in by_id]
async def get_item(pool: asyncpg.Pool, account_id: int, chat_id: int) -> DiscoverItem:
known = await by_ids(pool, account_id, [chat_id])
if known:
return known[0]
kind = await chat_kind(pool, account_id, chat_id)
return DiscoverItem(
chat_id=chat_id,
title=None,
username=None,
kind="private" if kind is ChatKind.DM else kind.value,
is_bot=False,
is_contact=False,
has_avatar=False,
message_count=0,
in_dialogs=False,
tracked=bool(await pool.fetchval(_IS_TRACKED, account_id, chat_id)),
)
async def chat_kind(pool: asyncpg.Pool, account_id: int, chat_id: int) -> ChatKind:
if chat_id > 0:
return ChatKind.DM
is_broadcast = await pool.fetchval(_IS_BROADCAST, account_id, chat_id)
return ChatKind.CHANNEL if is_broadcast else ChatKind.GROUP
+103 -9
View File
@@ -1,29 +1,123 @@
import asyncpg
from utils.read.models import MediaView
from utils.files import media_file_name
from utils.read.message_view import load_raw
from utils.read.models import MediaVersionView, MediaView
_MEDIA_COLS = (
"id, account_id, chat_id, message_id, kind, storage_key, file_size, "
"mime, ttl_seconds, downloaded, extracted_text, created_at"
MEDIA_COLS = (
"m.id, m.account_id, m.chat_id, m.message_id, m.kind, m.storage_key, "
"m.file_size, m.mime, m.ttl_seconds, m.downloaded, m.extracted_text, "
"m.created_at, src.original_name"
)
ORIGINAL_NAME_JOIN = """
LEFT JOIN LATERAL (
SELECT msg.raw->m.kind->>'file_name' AS original_name
FROM messages msg
WHERE msg.account_id = m.account_id AND msg.chat_id = m.chat_id
AND msg.message_id = m.message_id
ORDER BY msg.date DESC LIMIT 1
) src ON true
"""
_VERSION_COLS = "id, kind, storage_key, file_size, mime, observed_at"
_WEB_PAGE_MEDIA_KINDS = ("photo", "video", "animation", "document", "audio")
def media_view(row: asyncpg.Record) -> MediaView:
fields = dict(row)
original = fields.pop("original_name", None)
return MediaView(
**fields,
file_name=media_file_name(
fields["kind"], fields["mime"], fields["message_id"], original
),
)
async def _web_page_media_stub(
pool: asyncpg.Pool, account_id: int, chat_id: int, message_id: int
) -> MediaView | None:
row = await pool.fetchrow(
"SELECT date, raw FROM messages "
"WHERE account_id = $1 AND chat_id = $2 AND message_id = $3",
account_id,
chat_id,
message_id,
)
if row is None:
return None
web_page = load_raw(row["raw"]).get("web_page")
if not isinstance(web_page, dict):
return None
kind = next(
(k for k in _WEB_PAGE_MEDIA_KINDS if isinstance(web_page.get(k), dict)), None
)
if kind is None:
return None
obj = web_page[kind]
return MediaView(
id=0,
account_id=account_id,
chat_id=chat_id,
message_id=message_id,
kind=kind,
storage_key=None,
file_size=obj.get("file_size"),
mime=obj.get("mime_type"),
ttl_seconds=None,
downloaded=False,
extracted_text=None,
created_at=row["date"],
file_name=media_file_name(
kind, obj.get("mime_type"), message_id, obj.get("file_name")
),
)
async def get_media(pool: asyncpg.Pool, media_id: int) -> MediaView | None:
row = await pool.fetchrow(
f"SELECT {_MEDIA_COLS} FROM media WHERE id = $1", # noqa: S608
f"SELECT {MEDIA_COLS} FROM media m {ORIGINAL_NAME_JOIN} WHERE m.id = $1", # noqa: S608
media_id,
)
return MediaView(**dict(row)) if row else None
return media_view(row) if row else None
async def get_message_media(
pool: asyncpg.Pool, account_id: int, chat_id: int, message_id: int
) -> MediaView | None:
row = await pool.fetchrow(
f"SELECT {_MEDIA_COLS} FROM media " # noqa: S608
"WHERE account_id = $1 AND chat_id = $2 AND message_id = $3",
f"SELECT {MEDIA_COLS} FROM media m {ORIGINAL_NAME_JOIN} " # noqa: S608
"WHERE m.account_id = $1 AND m.chat_id = $2 AND m.message_id = $3",
account_id,
chat_id,
message_id,
)
return MediaView(**dict(row)) if row else None
if row is not None:
return media_view(row)
return await _web_page_media_stub(pool, account_id, chat_id, message_id)
async def get_media_versions(
pool: asyncpg.Pool, account_id: int, chat_id: int, message_id: int
) -> list[MediaVersionView]:
rows = await pool.fetch(
f"SELECT {_VERSION_COLS} FROM media_versions " # noqa: S608
"WHERE account_id = $1 AND chat_id = $2 AND message_id = $3 "
"ORDER BY observed_at",
account_id,
chat_id,
message_id,
)
return [MediaVersionView(**dict(row)) for row in rows]
async def get_media_version(
pool: asyncpg.Pool, version_id: int
) -> MediaVersionView | None:
row = await pool.fetchrow(
f"SELECT {_VERSION_COLS} FROM media_versions WHERE id = $1", # noqa: S608
version_id,
)
return MediaVersionView(**dict(row)) if row else None
+408
View File
@@ -0,0 +1,408 @@
import json
from datetime import datetime
from typing import Any
import asyncpg
from pydantic import ValidationError
from utils.files import media_file_name
from utils.read.models import (
ContactView,
EntityView,
ForwardView,
InlineButton,
LocationView,
MediaRef,
MessageView,
PollOption,
PollView,
ReactionCount,
ReplyView,
ServiceView,
StickerView,
WebPageView,
)
_MEDIA_KEYS = (
"photo",
"video",
"animation",
"voice",
"video_note",
"audio",
"document",
"sticker",
)
def load_raw(raw: str | None) -> dict[str, Any]:
if not raw:
return {}
try:
parsed = json.loads(raw)
except (ValueError, TypeError):
return {}
return parsed if isinstance(parsed, dict) else {}
def _enum(value: object) -> str | None:
if not isinstance(value, str):
return None
return value.rsplit(".", 1)[-1].lower()
def _parse_dt(value: object) -> datetime | None:
if not isinstance(value, str):
return None
try:
return datetime.fromisoformat(value)
except ValueError:
return None
def _peer_name(user: dict[str, Any]) -> str | None:
name = " ".join(
part for part in (user.get("first_name"), user.get("last_name")) if part
)
return name or user.get("username")
def _entities(raw: dict[str, Any]) -> list[EntityView]:
source = raw.get("entities") or raw.get("caption_entities") or []
if not isinstance(source, list):
return []
out: list[EntityView] = []
for item in source:
if not isinstance(item, dict):
continue
kind = _enum(item.get("type"))
offset = item.get("offset")
length = item.get("length")
if kind is None or not isinstance(offset, int) or not isinstance(length, int):
continue
custom = item.get("custom_emoji_id")
out.append(
EntityView(
type=kind,
offset=offset,
length=length,
url=item.get("url"),
custom_emoji_id=str(custom) if custom is not None else None,
language=item.get("language"),
)
)
return out
def _media_kind(message: dict[str, Any]) -> str | None:
kind = _enum(message.get("media"))
if kind:
return kind
for key in _MEDIA_KEYS:
if key in message:
return key
return None
def _reply(raw: dict[str, Any]) -> ReplyView | None:
reply = raw.get("reply_to_message")
reply_id = raw.get("reply_to_message_id")
if not isinstance(reply, dict):
return ReplyView(message_id=reply_id) if reply_id else None
sender = reply.get("from_user")
sender_chat = reply.get("sender_chat")
sender_id = None
sender_name = None
if isinstance(sender, dict):
sender_id = sender.get("id")
sender_name = _peer_name(sender)
elif isinstance(sender_chat, dict):
sender_id = sender_chat.get("id")
sender_name = sender_chat.get("title")
return ReplyView(
message_id=reply.get("id") or reply_id,
sender_id=sender_id,
sender_name=sender_name,
text=reply.get("text") or reply.get("caption"),
media_kind=_media_kind(reply),
)
def _forward(raw: dict[str, Any]) -> ForwardView | None:
origin = raw.get("forward_origin")
if not isinstance(origin, dict):
return None
tag = origin.get("_")
date = _parse_dt(origin.get("date"))
if tag == "MessageOriginUser":
user = origin.get("sender_user")
user = user if isinstance(user, dict) else {}
return ForwardView(
kind="user", from_id=user.get("id"), from_name=_peer_name(user), date=date
)
if tag == "MessageOriginChannel":
chat = origin.get("chat")
chat = chat if isinstance(chat, dict) else {}
return ForwardView(
kind="channel",
chat_id=chat.get("id"),
chat_title=chat.get("title"),
message_id=origin.get("message_id"),
signature=origin.get("author_signature"),
date=date,
)
return ForwardView(
kind="hidden", from_name=origin.get("sender_user_name"), date=date
)
def _reactions(raw: dict[str, Any]) -> list[ReactionCount]:
container = raw.get("reactions")
if not isinstance(container, dict):
return []
items = container.get("reactions")
if not isinstance(items, list):
return []
out: list[ReactionCount] = []
for item in items:
if not isinstance(item, dict):
continue
custom = item.get("custom_emoji_id")
out.append(
ReactionCount(
emoji=item.get("emoji"),
custom_emoji_id=str(custom) if custom is not None else None,
count=item.get("count") or 0,
chosen="chosen_order" in item,
)
)
return out
def _button_kind(button: dict[str, Any]) -> str:
if button.get("url"):
return "url"
if button.get("callback_data") is not None:
return "callback"
if "switch_inline_query" in button or "switch_inline_query_current_chat" in button:
return "switch"
return "other"
def _inline_buttons(raw: dict[str, Any]) -> list[list[InlineButton]]:
markup = raw.get("reply_markup")
if not isinstance(markup, dict):
return []
rows = markup.get("inline_keyboard")
if not isinstance(rows, list):
return []
out: list[list[InlineButton]] = []
for row in rows:
if not isinstance(row, list):
continue
buttons: list[InlineButton] = []
for button in row:
if not isinstance(button, dict):
continue
data = button.get("callback_data")
buttons.append(
InlineButton(
text=button.get("text") or "",
kind=_button_kind(button),
url=button.get("url"),
data=data if isinstance(data, str) else None,
)
)
if buttons:
out.append(buttons)
return out
def _web_page(raw: dict[str, Any]) -> WebPageView | None:
page = raw.get("web_page")
if not isinstance(page, dict) or not page.get("url"):
return None
return WebPageView(
url=page["url"],
display_url=page.get("display_url"),
type=page.get("type"),
site_name=page.get("site_name"),
title=page.get("title"),
description=page.get("description"),
has_photo="photo" in page,
)
def _text_of(value: dict[str, Any] | str | None) -> str | None:
if isinstance(value, dict):
text = value.get("text")
return text if isinstance(text, str) else None
return value if isinstance(value, str) else None
def _poll(raw: dict[str, Any]) -> PollView | None:
poll = raw.get("poll")
if not isinstance(poll, dict):
return None
raw_options = poll.get("options")
options: list[PollOption] = []
if isinstance(raw_options, list):
for option in raw_options:
if not isinstance(option, dict):
continue
options.append(
PollOption(
text=_text_of(option.get("text")) or "",
voter_count=option.get("voter_count") or 0,
vote_percentage=option.get("vote_percentage") or 0,
correct=option.get("is_correct"),
)
)
return PollView(
question=_text_of(poll.get("question")) or "",
options=options,
total_voter_count=poll.get("total_voter_count") or 0,
quiz=_enum(poll.get("type")) == "quiz",
closed=bool(poll.get("is_closed")),
multiple=bool(poll.get("allows_multiple_answers")),
anonymous=bool(poll.get("is_anonymous", True)),
)
def _contact(raw: dict[str, Any]) -> ContactView | None:
contact = raw.get("contact")
if not isinstance(contact, dict):
return None
return ContactView(
user_id=contact.get("user_id"),
first_name=contact.get("first_name"),
last_name=contact.get("last_name"),
phone_number=contact.get("phone_number"),
)
def _location(raw: dict[str, Any]) -> LocationView | None:
venue = raw.get("venue")
if isinstance(venue, dict):
point = venue.get("location")
point = point if isinstance(point, dict) else {}
return LocationView(
latitude=point.get("latitude"),
longitude=point.get("longitude"),
title=venue.get("title"),
address=venue.get("address"),
)
point = raw.get("location")
if not isinstance(point, dict):
return None
return LocationView(
latitude=point.get("latitude"), longitude=point.get("longitude")
)
def _service(raw: dict[str, Any]) -> ServiceView | None:
kind = _enum(raw.get("service"))
if kind is None:
return None
members = raw.get("new_chat_members") or raw.get("left_chat_member")
member_ids = None
if isinstance(members, list):
member_ids = [m["id"] for m in members if isinstance(m, dict) and "id" in m]
elif isinstance(members, dict) and "id" in members:
member_ids = [members["id"]]
pinned = raw.get("pinned_message")
call = raw.get("phone_call_ended")
return ServiceView(
kind=kind,
member_ids=member_ids,
pinned_message_id=pinned.get("id") if isinstance(pinned, dict) else None,
duration=call.get("duration") if isinstance(call, dict) else None,
)
def _sticker(raw: dict[str, Any]) -> StickerView | None:
sticker = raw.get("sticker")
if not isinstance(sticker, dict):
return None
return StickerView(
emoji=sticker.get("emoji"),
set_name=sticker.get("set_name"),
width=sticker.get("width"),
height=sticker.get("height"),
mime=sticker.get("mime_type"),
is_animated=bool(sticker.get("is_animated")),
is_video=bool(sticker.get("is_video")),
)
def media_ref_from(
message_id: int, raw: dict[str, Any], media_row: asyncpg.Record | None
) -> MediaRef | None:
kind = (media_row["kind"] if media_row else None) or _media_kind(raw)
if kind is None:
return None
obj = raw.get(kind)
obj = obj if isinstance(obj, dict) else {}
width = obj.get("width") or obj.get("length")
height = obj.get("height") or obj.get("length")
mime = (media_row["mime"] if media_row else None) or obj.get("mime_type")
return MediaRef(
message_id=message_id,
id=media_row["id"] if media_row else None,
kind=kind,
downloaded=bool(media_row["downloaded"]) if media_row else False,
width=width,
height=height,
duration=obj.get("duration"),
mime=mime,
file_size=(media_row["file_size"] if media_row else None)
or obj.get("file_size"),
ttl_seconds=media_row["ttl_seconds"] if media_row else None,
extracted_text=media_row["extracted_text"] if media_row else None,
file_name=media_file_name(kind, mime, message_id, obj.get("file_name")),
)
def _base_fields(row: asyncpg.Record) -> dict[str, Any]:
return {
"chat_id": row["chat_id"],
"message_id": row["message_id"],
"date": row["date"],
"sender_id": row["sender_id"],
"text": row["text"],
"has_media": row["has_media"],
"is_self_destruct": row["is_self_destruct"],
"edited_at": row["edited_at"],
"deleted_at": row["deleted_at"],
"media_group_id": row["media_group_id"],
}
def build_message_view(
row: asyncpg.Record, raw: dict[str, Any], media: list[MediaRef]
) -> MessageView:
base = _base_fields(row)
via_bot = raw.get("via_bot")
sticker = _sticker(raw)
try:
return MessageView(
**base,
entities=_entities(raw),
quote=_text_of(raw.get("quote")),
reply=_reply(raw),
forward=_forward(raw),
media=media,
reactions=_reactions(raw),
inline_buttons=_inline_buttons(raw),
web_page=_web_page(raw),
poll=_poll(raw),
contact=_contact(raw),
location=_location(raw),
service=_service(raw),
via_bot_id=via_bot.get("id") if isinstance(via_bot, dict) else None,
sticker=sticker,
is_sticker=sticker is not None,
is_animated_emoji=False,
)
except ValidationError:
return MessageView(**base, media=media)
+265
View File
@@ -16,11 +16,162 @@ class Page(BaseModel):
return min(self.limit, MAX_LIMIT)
class AccountView(BaseModel):
account_id: int
label: str | None
phone: str | None
tg_user_id: int | None
is_active: bool
device_model: str | None
class ChatListItem(BaseModel):
chat_id: int
title: str | None
kind: str
has_avatar: bool
is_bot: bool
is_contact: bool
is_broadcast: bool
message_count: int
last_date: datetime | None
last_text: str | None
last_sender_id: int | None
class DiscoverItem(BaseModel):
chat_id: int
title: str | None
username: str | None
kind: str
is_bot: bool
is_contact: bool
has_avatar: bool
message_count: int
in_dialogs: bool
tracked: bool
class EntityView(BaseModel):
type: str
offset: int
length: int
url: str | None = None
custom_emoji_id: str | None = None
language: str | None = None
class ReplyView(BaseModel):
message_id: int | None = None
sender_id: int | None = None
sender_name: str | None = None
text: str | None = None
media_kind: str | None = None
class ForwardView(BaseModel):
kind: str
from_id: int | None = None
from_name: str | None = None
chat_id: int | None = None
chat_title: str | None = None
message_id: int | None = None
date: datetime | None = None
signature: str | None = None
class MediaRef(BaseModel):
message_id: int
id: int | None = None
kind: str
downloaded: bool = False
width: int | None = None
height: int | None = None
duration: float | None = None
mime: str | None = None
file_size: int | None = None
ttl_seconds: int | None = None
extracted_text: str | None = None
file_name: str | None = None
class ReactionCount(BaseModel):
emoji: str | None = None
custom_emoji_id: str | None = None
count: int
chosen: bool = False
class InlineButton(BaseModel):
text: str
kind: str
url: str | None = None
data: str | None = None
class WebPageView(BaseModel):
url: str
display_url: str | None = None
type: str | None = None
site_name: str | None = None
title: str | None = None
description: str | None = None
has_photo: bool = False
class PollOption(BaseModel):
text: str
voter_count: int = 0
vote_percentage: int = 0
correct: bool | None = None
class PollView(BaseModel):
question: str
options: list[PollOption] = []
total_voter_count: int = 0
quiz: bool = False
closed: bool = False
multiple: bool = False
anonymous: bool = True
class ContactView(BaseModel):
user_id: int | None = None
first_name: str | None = None
last_name: str | None = None
phone_number: str | None = None
class LocationView(BaseModel):
latitude: float | None = None
longitude: float | None = None
title: str | None = None
address: str | None = None
class ServiceView(BaseModel):
kind: str
member_ids: list[int] | None = None
pinned_message_id: int | None = None
duration: int | None = None
class PinnedView(BaseModel):
message_id: int
text: str | None = None
media_kind: str | None = None
sender_name: str | None = None
class StickerView(BaseModel):
emoji: str | None = None
set_name: str | None = None
width: int | None = None
height: int | None = None
mime: str | None = None
is_animated: bool = False
is_video: bool = False
class MessageView(BaseModel):
@@ -33,6 +184,24 @@ class MessageView(BaseModel):
is_self_destruct: bool
edited_at: datetime | None
deleted_at: datetime | None
entities: list[EntityView] = []
quote: str | None = None
reply: ReplyView | None = None
forward: ForwardView | None = None
media_group_id: str | None = None
media: list[MediaRef] = []
reactions: list[ReactionCount] = []
inline_buttons: list[list[InlineButton]] = []
web_page: WebPageView | None = None
poll: PollView | None = None
contact: ContactView | None = None
location: LocationView | None = None
service: ServiceView | None = None
via_bot_id: int | None = None
sticker: StickerView | None = None
is_sticker: bool = False
is_animated_emoji: bool = False
read: bool = False
class MessageVersionView(BaseModel):
@@ -54,6 +223,36 @@ class MediaView(BaseModel):
downloaded: bool
extracted_text: str | None
created_at: datetime
file_name: str | None = None
class MediaVersionView(BaseModel):
id: int
kind: str
storage_key: str
file_size: int | None
mime: str | None
observed_at: datetime
class AvatarRef(BaseModel):
unique_id: str
storage_key: str | None
downloaded: bool
mime: str | None
class AvatarHistoryView(BaseModel):
unique_id: str
first_seen_at: datetime
downloaded: bool
class CustomEmojiRef(BaseModel):
storage_key: str | None
downloaded: bool
mime: str | None
kind: str | None
class CallbackView(BaseModel):
@@ -95,6 +294,20 @@ class PresenceHourly(BaseModel):
last_seen: datetime | None
class VolumeBucket(BaseModel):
bucket: datetime
total: int
outgoing: int
incoming: int
class ResponseStats(BaseModel):
mine_median_seconds: float | None
mine_count: int
their_median_seconds: float | None
their_count: int
class PeerView(BaseModel):
peer_id: int
first_name: str | None
@@ -103,6 +316,7 @@ class PeerView(BaseModel):
phone: str | None
photo_unique_id: str | None
is_deleted_account: bool
has_avatar: bool
updated_at: datetime
@@ -116,6 +330,27 @@ class PeerHistoryView(BaseModel):
is_deleted_account: bool
class ChatLinkView(BaseModel):
message_id: int
date: datetime | None
url: str
kind: str
web_url: str | None
web_title: str | None
web_site_name: str | None
class DayCount(BaseModel):
day: datetime
count: int
outgoing: int
class MessageAt(BaseModel):
message_id: int
date: datetime
class StoryView(BaseModel):
peer_id: int
story_id: int
@@ -158,3 +393,33 @@ class AlertView(BaseModel):
payload: dict[str, Any]
seen: bool
created_at: datetime
class FileShareView(BaseModel):
id: int
account_id: int
token: str
kind: str
file_name: str
url_name: str
mime: str | None
file_size: int | None
title: str | None
chat_id: int | None
message_id: int | None
peer_id: int | None
story_id: int | None
expires_at: datetime | None
max_downloads: int | None
download_count: int
last_download_at: datetime | None
revoked_at: datetime | None
created_at: datetime
class FileShareHitView(BaseModel):
ts: datetime
method: str
ip: str | None
user_agent: str | None
counted: bool
+36 -2
View File
@@ -2,13 +2,19 @@ import asyncpg
from utils.read.models import Page, PeerHistoryView, PeerView, StoryView
_PEER_COLS = (
"peer_id, first_name, last_name, username, phone, photo_unique_id, "
"is_deleted_account, updated_at, "
"EXISTS (SELECT 1 FROM avatars a WHERE a.account_id = peers.account_id "
"AND a.owner_id = peers.peer_id) AS has_avatar"
)
async def get_peer(
pool: asyncpg.Pool, account_id: int, peer_id: int
) -> PeerView | None:
row = await pool.fetchrow(
"SELECT peer_id, first_name, last_name, username, phone, "
"photo_unique_id, is_deleted_account, updated_at FROM peers "
f"SELECT {_PEER_COLS} FROM peers " # noqa: S608
"WHERE account_id = $1 AND peer_id = $2",
account_id,
peer_id,
@@ -16,6 +22,20 @@ async def get_peer(
return PeerView(**dict(row)) if row else None
async def get_peers(
pool: asyncpg.Pool, account_id: int, ids: list[int]
) -> list[PeerView]:
if not ids:
return []
rows = await pool.fetch(
f"SELECT {_PEER_COLS} FROM peers " # noqa: S608
"WHERE account_id = $1 AND peer_id = ANY($2)",
account_id,
ids,
)
return [PeerView(**dict(row)) for row in rows]
async def get_peer_history(
pool: asyncpg.Pool, account_id: int, peer_id: int
) -> list[PeerHistoryView]:
@@ -48,3 +68,17 @@ async def get_stories(
*params,
)
return [StoryView(**dict(row)) for row in rows]
async def get_story(
pool: asyncpg.Pool, account_id: int, peer_id: int, story_id: int
) -> StoryView | None:
row = await pool.fetchrow(
"SELECT peer_id, story_id, date, expire_date, caption, media_kind, "
"storage_key, downloaded, views, pinned, deleted FROM stories "
"WHERE account_id = $1 AND peer_id = $2 AND story_id = $3",
account_id,
peer_id,
story_id,
)
return StoryView(**dict(row)) if row else None
+32
View File
@@ -0,0 +1,32 @@
import asyncpg
from utils.read.message_view import _media_kind, _peer_name, load_raw
from utils.read.models import PinnedView
_PINNED_SQL = (
"SELECT raw FROM messages "
"WHERE account_id = $1 AND chat_id = $2 "
"AND raw->>'service' LIKE '%PINNED_MESSAGE%' "
"ORDER BY date DESC, message_id DESC LIMIT 1"
)
async def get_pinned(
pool: asyncpg.Pool, account_id: int, chat_id: int
) -> PinnedView | None:
row = await pool.fetchrow(_PINNED_SQL, account_id, chat_id)
if row is None:
return None
pinned = load_raw(row["raw"]).get("pinned_message")
if not isinstance(pinned, dict):
return None
message_id = pinned.get("id")
if message_id is None:
return None
sender = pinned.get("from_user")
return PinnedView(
message_id=message_id,
text=pinned.get("text") or pinned.get("caption"),
media_kind=_media_kind(pinned),
sender_name=_peer_name(sender) if isinstance(sender, dict) else None,
)
+13
View File
@@ -33,6 +33,19 @@ async def presence_history( # noqa: PLR0913
return [PresenceSample(**dict(row)) for row in rows]
async def current_presence(
pool: asyncpg.Pool, account_id: int, peer_id: int
) -> PresenceSample | None:
row = await pool.fetchrow(
"SELECT peer_id, ts, status, last_online_date, next_offline_date "
"FROM presence WHERE account_id = $1 AND peer_id = $2 "
"ORDER BY ts DESC LIMIT 1",
account_id,
peer_id,
)
return PresenceSample(**dict(row)) if row is not None else None
async def presence_hourly(
pool: asyncpg.Pool,
account_id: int,
+72
View File
@@ -0,0 +1,72 @@
from datetime import datetime, timedelta
import asyncpg
from utils.read.accounts import self_user_id
from utils.read.media import MEDIA_COLS, ORIGINAL_NAME_JOIN, media_view
from utils.read.models import ChatLinkView, DayCount, MediaView, MessageAt, Page
async def chat_media(
pool: asyncpg.Pool, account_id: int, chat_id: int, kinds: list[str], page: Page
) -> list[MediaView]:
rows = await pool.fetch(
f"SELECT {MEDIA_COLS} FROM media m {ORIGINAL_NAME_JOIN} " # noqa: S608
"WHERE m.account_id = $1 AND m.chat_id = $2 AND m.kind = ANY($3) "
"ORDER BY m.message_id DESC LIMIT $4 OFFSET $5",
account_id,
chat_id,
kinds,
page.capped_limit,
page.offset,
)
return [media_view(row) for row in rows]
async def chat_links(
pool: asyncpg.Pool, account_id: int, chat_id: int, page: Page
) -> list[ChatLinkView]:
rows = await pool.fetch(
"SELECT l.message_id, m.date, l.url, l.kind, l.web_url, "
"l.web_title, l.web_site_name FROM links l "
"LEFT JOIN messages m ON m.account_id = l.account_id "
"AND m.chat_id = l.chat_id AND m.message_id = l.message_id "
"WHERE l.account_id = $1 AND l.chat_id = $2 "
"ORDER BY l.message_id DESC, l.position LIMIT $3 OFFSET $4",
account_id,
chat_id,
page.capped_limit,
page.offset,
)
return [ChatLinkView(**dict(row)) for row in rows]
async def daily_counts(
pool: asyncpg.Pool, account_id: int, chat_id: int
) -> list[DayCount]:
self_id = await self_user_id(pool, account_id)
rows = await pool.fetch(
"SELECT date_trunc('day', date) AS day, count(*) AS count, "
"count(*) FILTER (WHERE sender_id = $3) AS outgoing FROM messages "
"WHERE account_id = $1 AND chat_id = $2 "
"GROUP BY day ORDER BY day",
account_id,
chat_id,
self_id,
)
return [DayCount(**dict(row)) for row in rows]
async def first_message_on_day(
pool: asyncpg.Pool, account_id: int, chat_id: int, day: datetime
) -> MessageAt | None:
row = await pool.fetchrow(
"SELECT message_id, date FROM messages "
"WHERE account_id = $1 AND chat_id = $2 AND date >= $3 AND date < $4 "
"ORDER BY date, message_id LIMIT 1",
account_id,
chat_id,
day,
day + timedelta(days=1),
)
return MessageAt(**dict(row)) if row else None
+10
View File
@@ -0,0 +1,10 @@
import asyncpg
async def read_up_to(pool: asyncpg.Pool, account_id: int, chat_id: int) -> int | None:
return await pool.fetchval(
"SELECT max(message_id) FROM read_receipts "
"WHERE account_id = $1 AND chat_id = $2 AND kind = 'read'",
account_id,
chat_id,
)
+207
View File
@@ -0,0 +1,207 @@
import secrets
from datetime import datetime
import asyncpg
from utils.files import url_slug
from utils.read.models import FileShareHitView, FileShareView, Page
_COLS = (
"id, account_id, token, kind, file_name, mime, file_size, title, "
"chat_id, message_id, peer_id, story_id, expires_at, max_downloads, "
"download_count, last_download_at, revoked_at, created_at"
)
_TOKEN_ALPHABET = "abcdefghijkmnpqrstuvwxyz23456789" # noqa: S105
_TOKEN_LENGTH = 10
_SERVE = """
SELECT id, kind, storage_key, file_name, mime, expires_at, max_downloads,
download_count, revoked_at
FROM file_shares WHERE token = $1
"""
_INSERT = (
"INSERT INTO file_shares " # noqa: S608
"(account_id, token, kind, storage_key, file_name, mime, file_size, title, "
"chat_id, message_id, peer_id, story_id, expires_at, max_downloads) "
"VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) "
f"RETURNING {_COLS}"
)
_ACTIVE = "revoked_at IS NULL AND (expires_at IS NULL OR expires_at > now())"
_LOOKUP = (
f"SELECT {_COLS} FROM file_shares " # noqa: S608
f"WHERE account_id = $1 AND storage_key = $2 AND {_ACTIVE} "
"AND (max_downloads IS NULL OR download_count < max_downloads) "
"ORDER BY created_at DESC LIMIT 1"
)
_CONSUME = """
UPDATE file_shares
SET download_count = download_count + 1, last_download_at = now()
WHERE id = $1 AND (max_downloads IS NULL OR download_count < max_downloads)
RETURNING download_count
"""
def new_token() -> str:
return "".join(secrets.choice(_TOKEN_ALPHABET) for _ in range(_TOKEN_LENGTH))
def _view(row: asyncpg.Record) -> FileShareView:
fields = dict(row)
return FileShareView(**fields, url_name=url_slug(fields["file_name"]))
async def list_shares(
pool: asyncpg.Pool, account_id: int, page: Page, *, active_only: bool = False
) -> list[FileShareView]:
where = "account_id = $1"
if active_only:
where += f" AND {_ACTIVE}"
rows = await pool.fetch(
f"SELECT {_COLS} FROM file_shares WHERE {where} " # noqa: S608
"ORDER BY created_at DESC LIMIT $2 OFFSET $3",
account_id,
page.capped_limit,
page.offset,
)
return [_view(row) for row in rows]
async def get_share(pool: asyncpg.Pool, share_id: int) -> FileShareView | None:
row = await pool.fetchrow(
f"SELECT {_COLS} FROM file_shares WHERE id = $1", # noqa: S608
share_id,
)
return _view(row) if row else None
async def find_active_share(
pool: asyncpg.Pool, account_id: int, storage_key: str
) -> FileShareView | None:
row = await pool.fetchrow(_LOOKUP, account_id, storage_key)
return _view(row) if row else None
async def create_share( # noqa: PLR0913
pool: asyncpg.Pool,
account_id: int,
kind: str,
storage_key: str,
file_name: str,
*,
mime: str | None = None,
file_size: int | None = None,
title: str | None = None,
chat_id: int | None = None,
message_id: int | None = None,
peer_id: int | None = None,
story_id: int | None = None,
expires_at: datetime | None = None,
max_downloads: int | None = None,
) -> FileShareView:
row = await pool.fetchrow(
_INSERT,
account_id,
new_token(),
kind,
storage_key,
file_name,
mime,
file_size,
title,
chat_id,
message_id,
peer_id,
story_id,
expires_at,
max_downloads,
)
return _view(row)
async def update_share(
pool: asyncpg.Pool,
share_id: int,
*,
expires_at: datetime | None,
max_downloads: int | None,
) -> FileShareView | None:
row = await pool.fetchrow(
"UPDATE file_shares SET expires_at = $2, max_downloads = $3 " # noqa: S608
f"WHERE id = $1 RETURNING {_COLS}",
share_id,
expires_at,
max_downloads,
)
return _view(row) if row else None
async def revoke_share(pool: asyncpg.Pool, share_id: int) -> FileShareView | None:
row = await pool.fetchrow(
"UPDATE file_shares SET revoked_at = now() " # noqa: S608
f"WHERE id = $1 AND revoked_at IS NULL RETURNING {_COLS}",
share_id,
)
if row is not None:
return _view(row)
return await get_share(pool, share_id)
async def delete_share(pool: asyncpg.Pool, share_id: int) -> bool:
result = await pool.execute("DELETE FROM file_shares WHERE id = $1", share_id)
return result.endswith("1")
async def rotate_token(pool: asyncpg.Pool, share_id: int) -> FileShareView | None:
row = await pool.fetchrow(
"UPDATE file_shares SET token = $2, revoked_at = NULL " # noqa: S608
f"WHERE id = $1 RETURNING {_COLS}",
share_id,
new_token(),
)
return _view(row) if row else None
async def share_by_token(pool: asyncpg.Pool, token: str) -> asyncpg.Record | None:
return await pool.fetchrow(_SERVE, token)
async def consume_download(pool: asyncpg.Pool, share_id: int) -> bool:
return await pool.fetchval(_CONSUME, share_id) is not None
async def record_hit( # noqa: PLR0913
pool: asyncpg.Pool,
share_id: int,
method: str,
ip: str | None,
user_agent: str | None,
*,
counted: bool,
) -> None:
await pool.execute(
"INSERT INTO file_share_hits (share_id, method, ip, user_agent, counted) "
"VALUES ($1, $2, $3, $4, $5)",
share_id,
method,
ip,
user_agent,
counted,
)
async def list_hits(
pool: asyncpg.Pool, share_id: int, page: Page
) -> list[FileShareHitView]:
rows = await pool.fetch(
"SELECT ts, method, ip, user_agent, counted FROM file_share_hits "
"WHERE share_id = $1 ORDER BY ts DESC LIMIT $2 OFFSET $3",
share_id,
page.capped_limit,
page.offset,
)
return [FileShareHitView(**dict(row)) for row in rows]
+4
View File
@@ -0,0 +1,4 @@
BEAVERGRAM_DOMAIN=beavergram.localhost
BEAVERGRAM_DEV_DOMAIN=dev.beavergram.localhost
CLOUDFLARE_API_TOKEN=
+2
View File
@@ -0,0 +1,2 @@
Caddyfile
.env
+19
View File
@@ -0,0 +1,19 @@
{
admin off
# acme_dns cloudflare {env.CLOUDFLARE_API_TOKEN}
log {
format console
}
servers {
trusted_proxies cloudflare
client_ip_headers Cf-Connecting-Ip
}
}
(compress) {
encode zstd gzip
}
import /etc/caddy/projects.d/*.caddy

Some files were not shown because too many files have changed in this diff Show More