diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 7428b93f..fa5e5f2a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -22,3 +22,49 @@ jobs: run: python -m pytest -q tests - name: Offline eval sanity run: python runtime/operations/scripts/offline_eval.py + + test-postgresql: + runs-on: ubuntu-latest + services: + postgres: + image: postgres:16-alpine + env: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: postgres + ports: + - 5432:5432 + options: >- + --health-cmd "pg_isready -U postgres" + --health-interval 10s + --health-timeout 5s + --health-retries 10 + env: + AIA_ASSISTANT_DB_BACKEND: postgresql + AIA_ASSISTANT_DATABASE_URL: postgresql://postgres:postgres@127.0.0.1:5432/oclaw + AIA_TEST_PG_URL: postgresql://postgres:postgres@127.0.0.1:5432/oclaw + OPS_WORKSPACE_ROOT: ${{ github.workspace }}/.ci_workspace + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: "3.11" + - name: Install dependencies + run: python -m pip install -r requirements.txt + - name: Create database and apply migrations + run: | + mkdir -p "$OPS_WORKSPACE_ROOT" + python - <<'PY' + import os + + import psycopg + + dsn = os.environ["AIA_ASSISTANT_DATABASE_URL"] + admin = dsn.rsplit("/", 1)[0] + "/postgres" + with psycopg.connect(admin, autocommit=True) as conn: + cur = conn.execute("SELECT 1 FROM pg_database WHERE datname = %s", ("oclaw",)) + if cur.fetchone() is None: + conn.execute("CREATE DATABASE oclaw") + PY + alembic upgrade head + - name: Run PostgreSQL-scoped tests + run: python -m pytest -q tests/test_database_backend.py tests/test_pg_compat.py tests/test_sa_app_settings.py tests/test_sa_auth_sessions.py tests/test_sa_chat_sessions.py tests/test_sa_chat_messages.py tests/test_sa_fork_trim_admin_sessions.py tests/test_sa_admin_user_stats_tool_health.py tests/test_sa_tool_log_trace.py tests/test_sa_prune_orphans.py tests/test_sa_ui_session_owner.py tests/test_sa_tenant_bind_code.py tests/test_sa_app_users.py tests/test_migrate_assistant_sqlite_to_pg.py diff --git a/_local/system.env.example b/_local/system.env.example index 34ac6a01..65df7ab3 100644 --- a/_local/system.env.example +++ b/_local/system.env.example @@ -46,6 +46,45 @@ AIA_PREWARM_INTERVAL_SECONDS=600 AIA_ASSISTANT_DB_PATH= OPS_ASSISTANT_DB_PATH= +# AIA_ASSISTANT_DB_BACKEND / OPS_ASSISTANT_DB_BACKEND 持久化后端:sqlite(默认)或 postgresql。 +# 设为 postgresql 时需配置 AIA_ASSISTANT_DATABASE_URL。首次连 PG 前在实例上建库(如 CREATE DATABASE oclaw),再应用 schema: +# 推荐:alembic upgrade head(或执行 svc/persistence/ddl/postgresql_bootstrap.sql),可选再运行 +# runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py 从 SQLite 迁数据。 +# +# 【Alembic】迁移脚本目录名为 assistant_migrations/(避免与 PyPI 包名 alembic 及 python -m alembic 冲突)。 +# 仓库根执行:python -m alembic upgrade head(或 alembic.exe -c alembic.ini upgrade head)。 +# 执行前请设置与本节相同的 PG 环境变量(至少 AIA_ASSISTANT_DB_BACKEND=postgresql 与 AIA_ASSISTANT_DATABASE_URL),以便 env.py 连上目标库。 +# +# 【割接到 PostgreSQL(生产/联调)建议顺序】 +# 1) 备份当前 SQLite 主库文件(db_path 指向的路径)。 +# 2) 在 PG 上 CREATE DATABASE …(如 oclaw),并对该库执行 alembic upgrade head。 +# 3) 运行导入(空库),任选其一(均依赖本机已安装 requirements.txt 含 psycopg、SQLAlchemy): +# · Linux / 无头设备:chmod +x runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh +# 后执行 ./runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh --dry-run 再去掉 --dry-run。 +# 或 ./runtime/operations/scripts/cutover_sqlite_to_postgresql.sh(含 SQLite 文件备份 + dry-run + 导入)。 +# · Windows:powershell -ExecutionPolicy Bypass -File runtime/operations/scripts/cutover_sqlite_to_postgresql.ps1 -LoadSystemEnv +# · 直接 Python(URL 可读环境变量或 _local/system.env + --load-system-env): +# python runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py --load-system-env --sqlite-from-db-path --dry-run +# 或显式:--sqlite <路径> [--pg-url ] +# 默认要求 PG 各表为空;若需叠加数据用 --allow-non-empty(自担风险)。 +# 4) 核对脚本末尾 verify 行数一致后,将 AIA_ASSISTANT_DB_BACKEND=postgresql 与 AIA_ASSISTANT_DATABASE_URL 写入 system.env / 部署环境并重启网关。 +# 【应用启动顺序】PostgreSQL 模式下须先对目标库执行 schema 迁移(alembic upgrade head),再启动网关;不要在空库上直接启动生产实例(缺表会失败)。数据割接时顺序为:备份 SQLite → 建库并迁移 → 导入脚本 → 再切换 env 并启动应用。 +AIA_ASSISTANT_DB_BACKEND= +OPS_ASSISTANT_DB_BACKEND= + +# AIA_ASSISTANT_DATABASE_URL / OPS_ASSISTANT_DATABASE_URL PostgreSQL 连接串。 +# 可用 libpq 形式:postgresql://user:pass@host:5432/dbname +# 也可用 SQLAlchemy 形式:postgresql+psycopg://user:pass@host:5432/dbname(与 SQLAlchemy URL 一致;原生 psycopg 连接前会自动去掉 +psycopg)。 +# 别名:AIA_ASSISTANT_PG_DSN / OPS_ASSISTANT_PG_DSN。 +# AIA_TEST_PG_URL 仅 pytest:若设置,tests/test_database_backend.py::test_get_assistant_store_postgresql_smoke 会连真实 PG; +# 可与 AIA_ASSISTANT_DATABASE_URL 填同一串(含 postgresql+psycopg:// 亦可)。 +AIA_ASSISTANT_DATABASE_URL= +OPS_ASSISTANT_DATABASE_URL= +AIA_TEST_PG_URL= + +# AIA_LOG_CHAT_MESSAGE_PERSIST 设为 1/true/on 时,每次成功写入 chat_message 后打一条 WARNING 日志(含 backend、session_id、message_id、role、event_type),便于在 gateway.err.log 里 grep「chat_message_persisted」;默认关闭。 +# AIA_LOG_CHAT_MESSAGE_PERSIST= + # AIA_ASSISTANT_MASTER_KEY 用于加密迁移密钥等;不设则部分迁移不可用。 # 【前端】管理后台「密钥迁移」相关界面会检测是否配置(提示文案)。 AIA_ASSISTANT_MASTER_KEY= diff --git a/alembic.ini b/alembic.ini new file mode 100644 index 00000000..05fc5723 --- /dev/null +++ b/alembic.ini @@ -0,0 +1,37 @@ +[alembic] +script_location = assistant_migrations +prepend_sys_path = . + +[loggers] +keys = root,sqlalchemy,alembic + +[handlers] +keys = console + +[formatters] +keys = generic + +[logger_root] +level = WARN +handlers = console +qualname = + +[logger_sqlalchemy] +level = WARN +handlers = +qualname = sqlalchemy.engine + +[logger_alembic] +level = INFO +handlers = +qualname = alembic + +[handler_console] +class = StreamHandler +args = (sys.stderr,) +level = NOTSET +formatter = generic + +[formatter_generic] +format = %(levelname)-5.5s [%(name)s] %(message)s +datefmt = %H:%M:%S diff --git a/assistant_migrations/env.py b/assistant_migrations/env.py new file mode 100644 index 00000000..0defba04 --- /dev/null +++ b/assistant_migrations/env.py @@ -0,0 +1,67 @@ +"""Alembic environment (assistant DB URL from env).""" + +from __future__ import annotations + +import os +import sys +from logging.config import fileConfig +from pathlib import Path + +from sqlalchemy import engine_from_config, pool, text +from sqlalchemy.engine import Connection + +from alembic import context + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +os.chdir(ROOT) + +from svc.config.bootstrap_env import load_system_env # noqa: E402 + +load_system_env() + +from svc.config.database import assistant_db_backend, assistant_sqlalchemy_url # noqa: E402 + +config = context.config +if config.config_file_name is not None: + fileConfig(config.config_file_name) + +target_metadata = None + + +def get_url() -> str: + return assistant_sqlalchemy_url() + + +def run_migrations_offline() -> None: + context.configure( + url=get_url(), + target_metadata=target_metadata, + literal_binds=True, + dialect_opts={"paramstyle": "named"}, + ) + with context.begin_transaction(): + context.run_migrations() + + +def run_migrations_online() -> None: + ini_section = config.get_section(config.config_ini_section) or {} + ini_section["sqlalchemy.url"] = get_url() + connectable = engine_from_config( + ini_section, + prefix="sqlalchemy.", + poolclass=pool.NullPool, + ) + + with connectable.connect() as connection: + context.configure(connection=connection, target_metadata=target_metadata) + with context.begin_transaction(): + context.run_migrations() + + +if context.is_offline_mode(): + run_migrations_offline() +else: + run_migrations_online() diff --git a/assistant_migrations/script.py.mako b/assistant_migrations/script.py.mako new file mode 100644 index 00000000..fbc4b07d --- /dev/null +++ b/assistant_migrations/script.py.mako @@ -0,0 +1,26 @@ +"""${message} + +Revision ID: ${up_revision} +Revises: ${down_revision | comma,n} +Create Date: ${create_date} + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +${imports if imports else ""} + +# revision identifiers, used by Alembic. +revision: str = ${repr(up_revision)} +down_revision: Union[str, None] = ${repr(down_revision)} +branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)} +depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)} + + +def upgrade() -> None: + ${upgrades if upgrades else "pass"} + + +def downgrade() -> None: + ${downgrades if downgrades else "pass"} diff --git a/assistant_migrations/versions/001_assistant_pg_initial.py b/assistant_migrations/versions/001_assistant_pg_initial.py new file mode 100644 index 00000000..ed73154a --- /dev/null +++ b/assistant_migrations/versions/001_assistant_pg_initial.py @@ -0,0 +1,34 @@ +"""Initial PostgreSQL schema for assistant store (from SQLite parity DDL).""" + +from __future__ import annotations + +from pathlib import Path + +from alembic import op +from sqlalchemy import text + +revision = "001_assistant_pg_initial" +down_revision = None +branch_labels = None +depends_on = None + + +def upgrade() -> None: + bind = op.get_bind() + if bind.dialect.name != "postgresql": + return + root = Path(__file__).resolve().parents[2] + sql_path = root / "svc" / "persistence" / "ddl" / "postgresql_bootstrap.sql" + raw = sql_path.read_text(encoding="utf-8") + for part in raw.split(";"): + stmt = part.strip() + if not stmt or stmt.startswith("--"): + continue + op.execute(text(stmt)) + + +def downgrade() -> None: + bind = op.get_bind() + if bind.dialect.name != "postgresql": + return + raise NotImplementedError("Assistant PG downgrade not supported; restore from backup.") diff --git a/docs/ENVIRONMENT_VARIABLES.md b/docs/ENVIRONMENT_VARIABLES.md index 54656c29..5da67401 100644 --- a/docs/ENVIRONMENT_VARIABLES.md +++ b/docs/ENVIRONMENT_VARIABLES.md @@ -599,6 +599,29 @@ ## 存储与迁移 +- `AIA_ASSISTANT_DB_BACKEND` / `OPS_ASSISTANT_DB_BACKEND` + - 默认:未设置(按 `sqlite` 处理) + - 取值:`sqlite`(默认)或 `postgresql`(大小写不敏感,亦接受 `pg` / `postgres` 等别名) + - 作用:主 assistant 持久化后端;设为 `postgresql` 时必须配置 `AIA_ASSISTANT_DATABASE_URL`(或 `OPS_*` / `*_PG_DSN` 别名) + - 生效:`oclaw/svc/config/database.py`, `oclaw/svc/persistence/assistant_store.py` + +- `AIA_ASSISTANT_DATABASE_URL` / `OPS_ASSISTANT_DATABASE_URL`(及 `AIA_ASSISTANT_PG_DSN` 等别名) + - 默认:空 + - 作用:PostgreSQL 连接串(`postgresql://…` 或 `postgresql+psycopg://…`) + - 生效:`oclaw/svc/config/database.py`, Alembic `assistant_migrations/env.py` + +- `AIA_TEST_PG_URL` + - 默认:空 + - 作用:仅测试用;设置后部分 pytest 会对真实 PostgreSQL 跑冒烟(如 `tests/test_database_backend.py`) + - 生效:对应测试模块 + +- **PostgreSQL 部署与启动顺序(摘要)** + 1. 在实例上 `CREATE DATABASE`(正式库名常用 `oclaw`),授予应用角色权限。 + 2. 配置与目标库一致的 `AIA_ASSISTANT_DB_BACKEND=postgresql` 与 `AIA_ASSISTANT_DATABASE_URL`,在**该库**上执行 **`alembic upgrade head`**(或等价执行 `svc/persistence/ddl/postgresql_bootstrap.sql`),完成建表。 + 3. 若从既有 SQLite 主库割接数据:在空表上运行 `runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py`(支持 `--load-system-env`、`--sqlite-from-db-path` 及从环境变量读取 PG URL;可先 `--dry-run`)。Linux 可用 `runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh`(仅导入)或 `cutover_sqlite_to_postgresql.sh`(备份 + dry-run + 导入);Windows 可用 `cutover_sqlite_to_postgresql.ps1`。 + 4. **再启动**网关/应用进程。不要在未迁移 schema 的空库上直接启动依赖表结构的生产服务。 + 5. 详细说明与 Alembic 在 Windows 上的注意项见仓库根 `_local/system.env.example` 第二节注释。 + - `AIA_ASSISTANT_DB_PATH` - 默认:`data/ai_ops.sqlite` - 作用:SQLite 路径 diff --git a/interfaces/admin/chat_api.py b/interfaces/admin/chat_api.py index 0528c551..140c84f4 100644 --- a/interfaces/admin/chat_api.py +++ b/interfaces/admin/chat_api.py @@ -31,10 +31,12 @@ from svc.files.file_attachments import ( from interfaces.ws.common import normalize_ws_attachments from svc.files.session_export import export_session_json, export_session_markdown from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.gateway import OclawGateway from runtime.plan_agent_v2.switch import v2_feature_enabled from runtime.types import StandardMessage, normalize_interaction_mode, normalize_requested_specialist from runtime.chat.history_tool_result_compact import compact_tool_results_in_session_history +from runtime.chat.persist_terminal_fallback import persist_assistant_text_if_turn_missing def _oclaw_config_path() -> Path: @@ -233,6 +235,23 @@ def _extract_manager_instruction_text(reasoning_content: str) -> str | None: return instr or None +def _instruction_from_dispatch_assignment_block(text: str) -> str | None: + """Extract manager instruction only from **comprehensive dispatch** reasoning blocks. + + Without this guard, any assistant ``reasoning_content`` that contains ``instruction:\\n`` followed + by the user's literal question (common in thinking traces) would register as a duplicate + ``instruction_text`` and :func:`_filter_internal_instruction_user_messages` would drop the real + ``user_text`` row — after a failed assistant persist the API could return **zero** messages on refresh. + """ + if not text or "instruction:\n" not in text: + return None + s = str(text) + tl = s.lower() + if "任务分配" not in s and "task assignment" not in tl: + return None + return _extract_manager_instruction_text(s) + + def _filter_internal_instruction_user_messages(msgs: list[Any]) -> list[Any]: """Hide legacy polluted user rows that equal manager dispatch instruction text.""" instruction_texts: set[str] = set() @@ -249,7 +268,7 @@ def _filter_internal_instruction_user_messages(msgs: list[Any]) -> list[Any]: if rc: candidates.append(rc) for text in candidates: - instr = _extract_manager_instruction_text(text) + instr = _instruction_from_dispatch_assignment_block(text) if instr: instruction_texts.add(instr) if not instruction_texts: @@ -362,6 +381,32 @@ def _resolve_chat_session(store: SqliteStore, ctx: dict[str, Any], session_id: s return store.get_session_for_user(session_id=session_id, tenant_id=tenant_id, user_id=user_id) +def _resolve_chat_session_allow_claim_orphan( + store: SqliteStore, ctx: dict[str, Any], session_id: str +): + """Like :func:`_resolve_chat_session`, but if ``chat_session`` exists with **no** ``ui_session_owner`` row + (common after SQLite→PG imports that only copied ``chat_session`` / ``chat_message``), attach the current + user once so ``get_session_for_user`` JOIN succeeds. Does **not** override an existing owner (wrong user). + """ + resolve = _resolve_chat_session + sess = resolve(store, ctx, session_id) + if sess is not None: + return sess + if _is_administrator_chat_viewer(ctx): + return None + sid = str(session_id or "").strip() + tenant_id = str(ctx.get("tenant_id") or "").strip() + user_id = str(ctx.get("user_id") or "").strip() + if not sid or not tenant_id or not user_id: + return None + if store.get_session(sid) is None: + return None + if store.get_ui_session_owner(session_id=sid) is not None: + return None + store.ensure_ui_session_owner(session_id=sid, tenant_id=tenant_id, user_id=user_id) + return resolve(store, ctx, session_id) + + def _effective_user_text(*, text: str, attachments: list[dict[str, Any]] | None, store: SqliteStore) -> str: """Streamlit 等价:仅有附件时也要落库一条用户消息,否则模型侧无输入。""" t = (text or "").strip() @@ -825,7 +870,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor offset: int = Query(default=0, ge=0), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") @@ -855,7 +900,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") @@ -881,18 +926,18 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") title = str(payload.get("title") or "").strip()[:_SESSION_TITLE_MAX_LEN] or ( "新会话" if _api_lang(store) == "zh" else "New Chat" ) store.rename_session(session_id, title) - s = _resolve_chat_session(store, ctx, session_id) + s = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) return { "ok": True, "session": { @@ -908,11 +953,11 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor session_id: str, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") if _is_administrator_chat_viewer(ctx): @@ -941,11 +986,11 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") last_id = store.get_last_message_id(session_id) @@ -976,11 +1021,11 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor format: str = Query(default="md", description="md or json"), authorization: str | None = Header(default=None), ) -> Response: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") fmt = (format or "md").strip().lower() @@ -1006,11 +1051,11 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor limit: int = Query(default=_CHAT_MSG_LIMIT, ge=1, le=20000), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") meta = store.get_session_messages_meta(session_id) @@ -1034,9 +1079,9 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor performance remains stable. It only touches `role=tool` chat_message rows. """ payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") limit_messages = int(payload.get("limit_messages") or 5000) @@ -1072,9 +1117,9 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor message_id: int, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") ok = store.delete_message(session_id=str(session_id), message_id=int(message_id)) @@ -1089,12 +1134,12 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor limit: int = Query(default=20, ge=1, le=200), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") after_iso = str(after or "").strip() @@ -1138,12 +1183,12 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor max_chars: int = Query(default=80_000, ge=1_000, le=200_000), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) tenant_id = str(ctx.get("tenant_id") or "") _ = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") wiki_root = _wiki_root_from_config() @@ -1177,11 +1222,11 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor session_id: str, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") s_mm, s_em = _resolve_session_dialog_chat_settings( @@ -1212,11 +1257,11 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") memory_mode = _normalize_memory_mode(payload) @@ -1254,7 +1299,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_user_mode_get( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") @@ -1274,7 +1319,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") @@ -1311,7 +1356,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor offset: int = Query(default=0, ge=0), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) tenant_id = str(ctx.get("tenant_id") or "") @@ -1333,7 +1378,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor offset: int = Query(default=0, ge=0), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) tenant_id = str(ctx.get("tenant_id") or "") @@ -1353,7 +1398,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor limit: int = Query(default=200, ge=20, le=1000), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) lang = _api_lang(store) labels = _dispatch_reason_labels_with_overrides(store) @@ -1392,7 +1437,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_get_ui_lang( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() _ = resolve_auth(store, authorization) return {"ok": True, "lang": _api_lang(store)} @@ -1402,7 +1447,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() _ = resolve_auth(store, authorization) lang = str(payload.get("lang") or "").strip().lower() if lang not in ("zh", "en"): @@ -1414,7 +1459,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_get_dispatch_reason_labels( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) raw = str(store.get_setting(_DISPATCH_REASON_LABELS_SETTING_KEY) or "").strip() @@ -1439,7 +1484,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: body = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) overrides = body.get("overrides") @@ -1465,7 +1510,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_get_specialist_flags( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) flags = _specialist_flags_with_overrides(store) @@ -1483,7 +1528,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: body = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) raw_flags = body.get("flags") @@ -1507,7 +1552,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor channel: str, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) ch = _normalize_channel_dispatch_channel(channel) @@ -1533,7 +1578,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: body = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) ch = _normalize_channel_dispatch_channel(channel) @@ -1553,7 +1598,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_get_attachment_limits( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) limits = _tabular_limits_from_oclaw_config() @@ -1565,7 +1610,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: body = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) raw = body.get("limits") @@ -1639,7 +1684,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_profile_get( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "").strip() user_id = str(ctx.get("user_id") or "").strip() @@ -1671,7 +1716,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "").strip() user_id = str(ctx.get("user_id") or "").strip() @@ -1702,7 +1747,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), file: UploadFile = File(...), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "").strip() user_id = str(ctx.get("user_id") or "").strip() @@ -1735,7 +1780,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor def api_chat_profile_avatar_delete( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "").strip() user_id = str(ctx.get("user_id") or "").strip() @@ -1749,9 +1794,9 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor session_id: str, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") sid = str(session_id) @@ -1766,7 +1811,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor attachment_id: str, authorization: str | None = Header(default=None), ) -> Response: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "").strip() user_id = str(ctx.get("user_id") or "").strip() @@ -1828,7 +1873,7 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor limit_messages: int = Query(default=50_000, ge=1, le=500_000), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_administrator_chat_viewer(ctx) tenant_id = str(ctx.get("tenant_id") or "").strip() @@ -1854,12 +1899,12 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") role = str(ctx.get("role") or "member") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") text_raw = str(payload.get("text") or "").strip() @@ -1971,6 +2016,13 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor specialist_executor_factory=specialist_factory, ) reply = gw_result.reply_text + persist_assistant_text_if_turn_missing( + store=store, + session_id=str(session_id), + turn_uuid=str(getattr(gw_result, "turn_uuid", "") or ""), + final_text=str(reply or ""), + log_prefix="admin_chat_http_text_fallback_persisted", + ) except GenerationInterrupted: msg = "已中断回答。" if lang == "zh" else "Response stopped." store.add_message(session_id=session_id, role="assistant", content=msg) @@ -1992,12 +2044,12 @@ def include_chat_routes(router: APIRouter, *, resolve_auth: Callable[[SqliteStor ) -> StreamingResponse: """Server-Sent Events: token deltas + progress + tool_ui, then done (assistant persisted by run_turn).""" payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) tenant_id = str(ctx.get("tenant_id") or "") user_id = str(ctx.get("user_id") or "") role = str(ctx.get("role") or "member") - sess = _resolve_chat_session(store, ctx, session_id) + sess = _resolve_chat_session_allow_claim_orphan(store, ctx, session_id) if not sess: raise HTTPException(status_code=404, detail="session_not_found") text_raw = str(payload.get("text") or "").strip() diff --git a/interfaces/admin/models_api.py b/interfaces/admin/models_api.py index f32258c0..60ea33a7 100644 --- a/interfaces/admin/models_api.py +++ b/interfaces/admin/models_api.py @@ -21,6 +21,7 @@ from runtime.agents.specialists import ( ) from svc.config.paths import db_path from runtime.orchestration.evaluation import eval_summary +from svc.persistence.assistant_store import get_assistant_store from svc.persistence.sqlite_store import ( LLM_BUILTIN_OLLAMA_PROFILE_ID, SqliteStore, @@ -169,7 +170,7 @@ def include_model_mgmt_routes( def api_models_state( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_permission(ctx, "admin:read") uid = str(ctx.get("user_id") or "").strip() @@ -224,7 +225,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) # 与能进入控制台一致:切换「当前选用」不写密钥,仅需读权限即可。 _require_permission(ctx, "admin:read") @@ -242,7 +243,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) profiles = store.list_llm_profiles(visible_only=True, **_models_list_kwargs(ctx)) @@ -268,7 +269,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) profiles = store.list_llm_profiles(visible_only=True, **_models_list_kwargs(ctx)) @@ -285,7 +286,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) profiles = store.list_llm_profiles(visible_only=True, **_models_list_kwargs(ctx)) @@ -311,7 +312,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") visible = _truthy(payload.get("chat_model_selector_visible"), default=True) @@ -324,7 +325,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) name = str(payload.get("name") or "").strip() or "新配置" @@ -354,7 +355,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) pid = str(profile_id or "").strip() @@ -401,7 +402,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) pid = str(profile_id or "").strip() @@ -428,7 +429,7 @@ def include_model_mgmt_routes( profile_id: str, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_models_mutate(ctx) pid = str(profile_id or "").strip() @@ -446,7 +447,7 @@ def include_model_mgmt_routes( def api_models_members( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -481,7 +482,7 @@ def include_model_mgmt_routes( profile_id: str = Query(...), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -497,7 +498,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -524,7 +525,7 @@ def include_model_mgmt_routes( profile_id: str = Query(...), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -539,7 +540,7 @@ def include_model_mgmt_routes( profile_id: str = Query(..., description="llm_profile id"), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -555,7 +556,7 @@ def include_model_mgmt_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -589,7 +590,7 @@ def include_model_mgmt_routes( user_id: str = Query(...), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_grant_manager(ctx) tid = str(ctx.get("tenant_id") or "").strip() @@ -606,7 +607,7 @@ def include_model_mgmt_routes( limit_logs: int = Query(default=100, ge=1, le=500), limit_summary: int = Query(default=500, ge=1, le=5000), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_permission(ctx, "admin:read") summary = eval_summary(store, limit=limit_summary) @@ -630,7 +631,7 @@ def include_model_mgmt_routes( format: str = Query(default="csv", description="csv or json"), limit: int = Query(default=100_000, ge=1, le=200_000), ) -> Response: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_permission(ctx, "admin:read") fmt = str(format or "csv").strip().lower() diff --git a/interfaces/admin/routes.py b/interfaces/admin/routes.py index 56520b5e..20fc9859 100644 --- a/interfaces/admin/routes.py +++ b/interfaces/admin/routes.py @@ -34,6 +34,7 @@ from runtime.orchestration.vector_store import read_vector_memory_runtime from svc.config.paths import PROJECT_ROOT, db_path from svc.config.passwords import load_expected_password from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.agents.specialists import discover_specialist_ids, parse_agent_profile_bindings from runtime.tools.mcp.installer import ( _safe_server_id, @@ -442,7 +443,7 @@ def build_admin_router() -> APIRouter: def api_internal_tools_reload( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") @@ -467,7 +468,7 @@ def build_admin_router() -> APIRouter: def api_tools_exposure_trace_setting_get( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") raw = str(store.get_setting("AIA_TRACE_TOOL_EXPOSURE_PLAN") or "").strip().lower() @@ -479,7 +480,7 @@ def build_admin_router() -> APIRouter: payload: dict[str, Any] | None = Body(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") payload = payload or {} @@ -501,7 +502,7 @@ def build_admin_router() -> APIRouter: role: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") from runtime.tools.exposure_plan import build_internal_tool_specs @@ -539,7 +540,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: """One-shot diagnostic summary for role-based tool exposure.""" - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") @@ -604,7 +605,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: """Preview the final tools injected to LLM for a role (internal + MCP + wire policy).""" - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") from runtime.tools.exposure_plan import build_llm_tools_plan @@ -726,7 +727,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/channels") def api_channels(authorization: str | None = Header(default=None)) -> dict[str, Any]: - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:read") reg = build_channel_registry() items = [{"name": k, "type": reg[k].__class__.__name__} for k in sorted(reg.keys())] @@ -734,7 +735,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/stack/status") def api_stack_status(authorization: str | None = Header(default=None)) -> dict[str, Any]: - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:read") items = [] for s in status_services(): @@ -786,7 +787,7 @@ def build_admin_router() -> APIRouter: @router.post("/admin/api/stack/up") def api_stack_up(channel: str = "wecom", authorization: str | None = Header(default=None)) -> dict[str, Any]: - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:runtime:write") # Use same defaults as CLI; this starts detached processes and writes runtime state. import argparse @@ -807,7 +808,7 @@ def build_admin_router() -> APIRouter: @router.post("/admin/api/stack/down") def api_stack_down(authorization: str | None = Header(default=None)) -> dict[str, Any]: - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:runtime:write") import argparse @@ -884,14 +885,14 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/runtime/anomalies") def api_runtime_anomalies(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") return _collect_runtime_anomalies(store) @router.post("/admin/api/runtime/cleanup") def api_runtime_cleanup(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:runtime:write") killed: list[dict[str, Any]] = [] @@ -928,7 +929,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/runtime/prewarm/status") def api_runtime_prewarm_status(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") return runtime_prewarm_status(store=store) @@ -938,7 +939,7 @@ def build_admin_router() -> APIRouter: payload: dict[str, Any] | None = Body(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:runtime:write") body = payload or {} @@ -950,7 +951,7 @@ def build_admin_router() -> APIRouter: def _run() -> None: try: - _ = run_runtime_prewarm(reason=reason, store=SqliteStore(db_path())) + _ = run_runtime_prewarm(reason=reason, store=get_assistant_store()) except Exception: pass @@ -963,14 +964,14 @@ def build_admin_router() -> APIRouter: role: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") return runtime_prewarm_prompts_snapshot(store=store, role=str(role or "").strip().lower() or None) @router.get("/admin/api/runtime/scan-artifacts") def api_runtime_scan_artifacts(authorization: str | None = Header(default=None)) -> dict[str, Any]: - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:read") root = (PROJECT_ROOT / "runtime" / "data" / "scan").resolve() allowed_prefixes = ("history_entries_", "state_scan_") @@ -995,7 +996,7 @@ def build_admin_router() -> APIRouter: @router.post("/admin/api/runtime/scan-artifacts/cleanup") def api_runtime_scan_artifacts_cleanup(authorization: str | None = Header(default=None)) -> dict[str, Any]: - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:runtime:write") root = (PROJECT_ROOT / "runtime" / "data" / "scan").resolve() allowed_prefixes = ("history_entries_", "state_scan_") @@ -1017,7 +1018,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: body = payload or {} - ctx = _resolve_auth(SqliteStore(db_path()), authorization) + ctx = _resolve_auth(get_assistant_store(), authorization) _require_permission(ctx, "admin:runtime:write") root = (PROJECT_ROOT / "runtime" / "data" / "scan").resolve() try: @@ -1068,7 +1069,7 @@ def build_admin_router() -> APIRouter: scope: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:read") rows = store.list_tenants(limit=500) @@ -1083,7 +1084,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") name = str(payload.get("name") or "").strip() or "Team" @@ -1096,7 +1097,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1128,7 +1129,7 @@ def build_admin_router() -> APIRouter: user_id: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:read") _require_tenant_scope(ctx, tenant_id) @@ -1152,7 +1153,7 @@ def build_admin_router() -> APIRouter: include_inactive: int = Query(default=1), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:read") _require_tenant_scope(ctx, tenant_id) @@ -1172,7 +1173,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1234,7 +1235,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1269,7 +1270,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=1000), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:read") _require_tenant_scope(ctx, tenant_id) @@ -1288,7 +1289,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") tenant_id = str(payload.get("tenant_id") or "").strip() or str(ctx.get("tenant_id") or "").strip() @@ -1327,7 +1328,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1366,7 +1367,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:delete") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1435,7 +1436,7 @@ def build_admin_router() -> APIRouter: user_id: str = Query(default=""), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) rmode = _workspace_path_policy_read_mode(ctx) if rmode is None: @@ -1468,7 +1469,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) wmode = _workspace_path_policy_write_mode(ctx) if wmode is None: @@ -1518,7 +1519,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:delete") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1557,7 +1558,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/bind-codes") def api_bind_codes(tenant_id: str, authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:read") _require_tenant_scope(ctx, tenant_id) @@ -1570,7 +1571,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1588,7 +1589,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") tenant_id = str(payload.get("tenant_id") or "").strip() @@ -1608,7 +1609,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/wecom/config") def api_wecom_config(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") return { @@ -1625,7 +1626,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/wecom/health") def api_wecom_health(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") svc = next((s for s in status_services() if str(s.name) == "channel:wecom"), None) @@ -1672,7 +1673,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") store.set_setting("wecom_mode", "bot_api") @@ -1695,7 +1696,7 @@ def build_admin_router() -> APIRouter: @router.post("/admin/api/wecom/unbind") def api_wecom_unbind(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") with store._connect() as conn: @@ -1722,7 +1723,7 @@ def build_admin_router() -> APIRouter: session_id: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") rows = store.list_agent_audit_logs(limit=200, session_id=session_id) @@ -1734,7 +1735,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=80), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") rows = store.list_session_tool_health(session_id=session_id, limit=limit) @@ -1746,7 +1747,7 @@ def build_admin_router() -> APIRouter: session_id: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") if not session_id: @@ -1763,7 +1764,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=80), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") tenant_id = str(ctx.get("tenant_id") or "") @@ -1794,7 +1795,7 @@ def build_admin_router() -> APIRouter: include_attempts: int = Query(default=1), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") tenant_id = str(ctx.get("tenant_id") or "") @@ -1840,7 +1841,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: """Return a best-effort replay bundle for a single turn (trace_id).""" - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") sid = str(session_id or "").strip() @@ -1940,7 +1941,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/plugins") def api_plugins(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") rows = store.list_tool_plugins() @@ -1948,7 +1949,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/tool-policy") def api_tool_policy(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") # Oclaw takeover: legacy tool-policy switches are disconnected (kept in DB for later). @@ -2047,7 +2048,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") # Oclaw takeover: legacy tool-policy switches are disconnected (kept in DB for later). @@ -2240,7 +2241,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/mcp/servers") def api_mcp_servers(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") rows = McpRegistry(store).list_servers(enabled_only=False) @@ -2253,7 +2254,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/mcp/export") def api_mcp_export(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") return { @@ -2267,7 +2268,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=20), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") return {"ok": True, "items": store.list_mcp_install_failure_summary(limit=limit)} @@ -2278,7 +2279,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=200), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") return { @@ -2292,7 +2293,7 @@ def build_admin_router() -> APIRouter: role: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") from svc.llm.tool_wire_policy import build_tool_wire_snapshot @@ -2307,7 +2308,7 @@ def build_admin_router() -> APIRouter: from svc.llm.tool_wire_policy import SETTINGS_KEY_ADMIN_CONFIG, load_merged_admin_config payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") raw = store.get_setting(SETTINGS_KEY_ADMIN_CONFIG) @@ -2359,7 +2360,7 @@ def build_admin_router() -> APIRouter: from svc.llm.tool_wire_policy import SETTINGS_KEY_ROLE_MODE_BY_ROLE payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") role = str(payload.get("role") or "").strip().lower() @@ -2398,7 +2399,7 @@ def build_admin_router() -> APIRouter: ) -> dict[str, Any]: from svc.llm.tool_wire_policy import SETTINGS_KEY_PENALTY_STATE, SETTINGS_KEY_PENALTY_STATE_BY_ROLE - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") r = str(role or "").strip().lower() @@ -2440,7 +2441,7 @@ def build_admin_router() -> APIRouter: ) payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") role = str(payload.get("role") or "").strip().lower() @@ -2511,7 +2512,7 @@ def build_admin_router() -> APIRouter: ) payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") try: @@ -2570,7 +2571,7 @@ def build_admin_router() -> APIRouter: per_source_limit: int = Query(default=6), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") query = str(q or "").strip() @@ -2585,7 +2586,7 @@ def build_admin_router() -> APIRouter: refresh: int = Query(default=0), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") items = trending_mcp_market(force_refresh=bool(int(refresh or 0)), per_source_limit=per_source_limit) @@ -2593,14 +2594,14 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/mcp/dependencies") def api_mcp_dependencies(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") return {"ok": True, "items": detect_local_dependencies()} @router.get("/admin/api/mcp/binding") def api_mcp_binding(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") available = _ordered_mcp_roles() @@ -2627,7 +2628,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") available = _ordered_mcp_roles() @@ -2661,7 +2662,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/mcp/specialists") def api_mcp_specialists(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") available = _ordered_mcp_roles() @@ -2679,7 +2680,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") available = _ordered_mcp_roles() @@ -2705,7 +2706,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/experts") def api_experts_list(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") items = [_serialize_expert_row(x) for x in list_experts()] @@ -2717,7 +2718,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") expert_id = normalize_expert_id(payload.get("id")) @@ -2755,7 +2756,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") files_raw = payload.get("files") if isinstance(payload.get("files"), dict) else {} @@ -2788,7 +2789,7 @@ def build_admin_router() -> APIRouter: @router.delete("/admin/api/experts/{expert_id}") def api_experts_delete(expert_id: str, authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") eid = normalize_expert_id(expert_id) @@ -2814,7 +2815,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") source_type = str(payload.get("source_type") or "").strip().lower() @@ -2882,7 +2883,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") server_id = str(payload.get("server_id") or "").strip() @@ -2952,7 +2953,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") server_id = str(payload.get("server_id") or "").strip() @@ -3073,7 +3074,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") server_id = str(payload.get("server_id") or "").strip() @@ -3133,7 +3134,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") server_id = str(payload.get("server_id") or "").strip() @@ -3163,7 +3164,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") source_type = str(payload.get("source_type") or "").strip().lower() @@ -3192,7 +3193,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") server_id = str(payload.get("server_id") or "").strip() @@ -3217,7 +3218,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") apply_gateway_mcp_env_to_os() @@ -3268,7 +3269,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") apply_gateway_mcp_env_to_os() @@ -3327,7 +3328,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") apply_gateway_mcp_env_to_os() @@ -3347,7 +3348,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") enabled_only = bool(payload.get("enabled_only", True)) @@ -3372,7 +3373,7 @@ def build_admin_router() -> APIRouter: run the same health + tools/list + replace flow as ``check-all`` (only those servers). """ payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") apply_gateway_mcp_env_to_os() @@ -3426,7 +3427,7 @@ def build_admin_router() -> APIRouter: def api_mcp_e2e_check( authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") apply_gateway_mcp_env_to_os() @@ -3520,7 +3521,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=100), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") if tenant_id: @@ -3535,7 +3536,7 @@ def build_admin_router() -> APIRouter: limit: int = Query(default=300), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") if tenant_id: @@ -3567,7 +3568,7 @@ def build_admin_router() -> APIRouter: offset: int = Query(default=0), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") if tenant_id: @@ -3582,7 +3583,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:memory:write") memory_id = str(payload.get("memory_id") or "").strip() @@ -3597,7 +3598,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:memory:write") tenant_id = str(payload.get("tenant_id") or "").strip() or None @@ -3635,7 +3636,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/memory/config") def api_memory_config(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") runtime = read_vector_memory_runtime(store) @@ -3666,7 +3667,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:memory:write") store.set_setting("MEMORY_VECTOR_ENABLED", "1" if str(payload.get("enabled") or "").lower() in ("1", "true", "yes", "on") else "0") @@ -3704,7 +3705,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:memory:write") try: @@ -3717,7 +3718,7 @@ def build_admin_router() -> APIRouter: @router.post("/admin/api/secrets/migrate") def api_secrets_migrate(authorization: str | None = Header(default=None)) -> dict[str, Any]: """Migrate legacy b64 secrets to fernet (requires AIA_ASSISTANT_MASTER_KEY).""" - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:tenant:write") try: @@ -3738,7 +3739,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/secrets/status") def api_secrets_status(authorization: str | None = Header(default=None)) -> dict[str, Any]: """Expose legacy secret stats for admin UI warnings.""" - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:read") stats = store.legacy_secret_stats() @@ -3749,14 +3750,14 @@ def build_admin_router() -> APIRouter: @router.post("/admin/api/auth/bootstrap") def api_auth_bootstrap() -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() _ensure_admin_bootstrap(store) return {"ok": True} @router.post("/admin/api/auth/login") def api_auth_login(payload: dict[str, Any] | None = Body(default=None)) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() _ensure_admin_bootstrap(store) tenant_id = str(payload.get("tenant_id") or "").strip() if not tenant_id: @@ -3820,13 +3821,13 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/auth/me") def api_auth_me(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) return {"ok": True, "session": ctx} @router.post("/admin/api/auth/logout") def api_auth_logout(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() token = _extract_bearer(authorization) if token: store.revoke_auth_session(session_token_hash=_sha256_hex(token)) @@ -3838,7 +3839,7 @@ def build_admin_router() -> APIRouter: authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_ops_ai_auth(store, authorization) _require_permission(ctx, "admin:read") @@ -4045,7 +4046,7 @@ def build_admin_router() -> APIRouter: offset: int = Query(default=0), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_ops_ai_auth(store, authorization) _require_permission(ctx, "admin:read") lim = max(1, min(int(limit), 200)) @@ -4068,7 +4069,7 @@ def build_admin_router() -> APIRouter: @router.get("/admin/api/ops-ai/health") def api_ops_ai_health(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_ops_ai_auth(store, authorization) _require_permission(ctx, "admin:read") return {"ok": True, "service": "oclaw", "component": "ops-ai", "status": "ok", "caller": str(ctx.get("username") or "")} @@ -4082,7 +4083,7 @@ def build_admin_router() -> APIRouter: status: str | None = Query(default=None), authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = _resolve_auth(store, authorization) _require_permission(ctx, "admin:user:write") a = str(action or "").strip() diff --git a/interfaces/admin/skills_api.py b/interfaces/admin/skills_api.py index 185baad7..dbf8abb8 100644 --- a/interfaces/admin/skills_api.py +++ b/interfaces/admin/skills_api.py @@ -36,6 +36,7 @@ from runtime.skills_market import get_market_adapter, normalize_skill_market_pro from runtime.tools.skills_runtime.subprocess_exec import run_skill_runtime_entry from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store _SKILL_MARKET_PROVIDER_KEY = "AIA_SKILL_MARKET_PROVIDER" @@ -85,7 +86,7 @@ def include_skill_routes( @sk.get("") def api_skills_list(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) items = list_skills_with_status(store=store) @@ -93,7 +94,7 @@ def include_skill_routes( @sk.get("/mode") def api_skills_mode_get(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) raw_prompt = str(store.get_setting("AIA_SKILLS_PROMPT_IN_SYSTEM") or "").strip().lower() @@ -114,7 +115,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) if "prompt_in_system" in payload: @@ -153,7 +154,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) source_dir = str(payload.get("source_dir") or "").strip() @@ -187,7 +188,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) archive_url = str(payload.get("archive_url") or "").strip() @@ -221,7 +222,7 @@ def include_skill_routes( limit: int | None = None, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) query = str(q or "").strip() @@ -236,7 +237,7 @@ def include_skill_routes( slug: str | None = None, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) s = str(slug or "").strip() @@ -252,7 +253,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) s = str(payload.get("slug") or "").strip() @@ -309,7 +310,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -346,7 +347,7 @@ def include_skill_routes( @sk.get("/binding") def api_skills_binding_get(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_tenant_write(ctx) roles, mapping, _valid = _normalized_skill_binding(store) @@ -368,7 +369,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_tenant_write(ctx) if "enabled" in payload: @@ -403,7 +404,7 @@ def include_skill_routes( @sk.get("/effective") def api_skills_effective(authorization: str | None = Header(default=None)) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) roles, mapping, _valid = _normalized_skill_binding(store) @@ -494,7 +495,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -538,7 +539,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -554,7 +555,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -570,7 +571,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -604,7 +605,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) auto_name = str(payload.get("name") or "") @@ -642,7 +643,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) source = str(payload.get("source") or "").strip().lower() @@ -706,7 +707,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -736,7 +737,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: _payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) items = list_skills_with_status(store=store) @@ -798,7 +799,7 @@ def include_skill_routes( authorization: str | None = Header(default=None), ) -> dict[str, Any]: payload = payload or {} - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) name = str(payload.get("name") or "").strip() @@ -850,7 +851,7 @@ def include_skill_routes( include_execution: bool = False, authorization: str | None = Header(default=None), ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() ctx = resolve_auth(store, authorization) _require_admin(ctx) manifests = discover_workspace_skill_manifests() diff --git a/interfaces/admin/static/chat.js b/interfaces/admin/static/chat.js index aca8f1b9..8808b99f 100644 --- a/interfaces/admin/static/chat.js +++ b/interfaces/admin/static/chat.js @@ -688,6 +688,28 @@ function _buildRenderRows(msgs) { return rows; } +/** True when persisted history does not show a completed assistant reply for the latest turn (e.g. ends with user only). */ +function _needsWsTextFallbackFromRenderRows(renderRows) { + const rows = Array.isArray(renderRows) ? renderRows : []; + if (!rows.length) return true; + const last = rows[rows.length - 1]; + const lr = String((last && last.role) || "").toLowerCase(); + if (lr === "user") return true; + if (lr !== "assistant") return true; + const items = last && Array.isArray(last._items) ? last._items : []; + for (const it of items) { + const k = String((it && it.kind) || "").toLowerCase(); + if (k === "assistant_text" && String((it && it.text) || "").trim()) return false; + if (k === "reasoning" && String((it && it.text) || "").trim()) return false; + if (k === "tool_result") { + if (String((it && it.text) || "").trim()) return false; + if (Array.isArray(it.attachments) && it.attachments.length) return false; + } + } + if (String((last && last.content) || "").trim()) return false; + return true; +} + function t(key, vars) { let s = (I18N[currentLang] && I18N[currentLang][key]) || (I18N.en && I18N.en[key]) || key; if (vars && typeof s === "string") { @@ -3739,25 +3761,35 @@ async function renderChatUi() { }, }; - const loadMessagesForActive = async () => { + const loadMessagesForActive = async (opts = {}) => { loadMessagesForActive._rid = (loadMessagesForActive._rid || 0) + 1; const rid = loadMessagesForActive._rid; shouldFollowMessages = true; if (!activeId) { + loadMessagesForActive._needsWsTextFallback = true; messagesEl.innerHTML = ""; statusBar.textContent = sessions.length ? "" : t("chat.noSessions"); if (!sessions.length) messagesEl.appendChild(el("div", { class: "muted", text: t("chat.empty") })); - return; + return 0; } statusBar.textContent = t("chat.loading"); + loadMessagesForActive._needsWsTextFallback = true; try { const resp = await apiGet( `/admin/api/chat/sessions/${encodeURIComponent(activeId)}/messages?limit=${CHAT_MESSAGES_FETCH_LIMIT}`, ); - if (rid !== loadMessagesForActive._rid) return; + if (rid !== loadMessagesForActive._rid) return 0; const msgs = Array.isArray(resp.messages) ? resp.messages : []; const renderRows = _buildRenderRows(msgs); + loadMessagesForActive._needsWsTextFallback = _needsWsTextFallbackFromRenderRows(renderRows); const total = intOr(resp.message_count, msgs.length); + // End-of-turn hydrate: if the server returns nothing renderable yet (PG commit lag, transient API + // glitch) but the caller still has a live stream worth keeping, do not wipe messagesEl — that + // caused "stream flashes then entire dialog is empty" when reasoning/tool-output mode reloads. + if (!renderRows.length && opts.keepDomIfNoHistoryRows) { + statusBar.textContent = ""; + return 0; + } messagesEl.innerHTML = ""; if (total > msgs.length) { messagesEl.appendChild( @@ -3776,10 +3808,13 @@ async function renderChatUi() { } statusBar.textContent = ""; scrollMessagesToBottom(true); + return renderRows.length; } catch (e) { // Keep existing message list on reload failure (e.g. toggle reload races), // so users don't perceive "messages disappeared". + loadMessagesForActive._needsWsTextFallback = true; statusBar.textContent = `${t("chat.error")}: ${String(e)}`; + return -1; } }; @@ -4847,25 +4882,39 @@ ${autoLimit ? `
auto-added claus renderStreamComposite(); _markStreamTerminal("end", t("chat.status.end")); ok = true; - } else if (adminChatShowToolOutput) { - // WS final message may only contain the last assistant_text snapshot and - // omit persisted reasoning rows. Hydrate from history to avoid - // end-of-turn "reasoning disappears until refresh". - try { - if (streamRow && streamRow.parentNode) streamRow.remove(); - } catch (_) {} - await loadMessagesForActive(); + } else if (adminChatShowToolOutput || sawStreamToolRefAttachments) { + // Hydrate from history for reasoning / tool panels / ref attachments. If the DB has not + // yet persisted the assistant reply (or only has the user row), n>0 used to skip + // appendFinalAssistant and removed the stream bubble — leaving an empty pane like the + // user screenshot. Use _needsWsTextFallback when WS still holds usable text. + const hadStream = + hasRealStreamText || + (Array.isArray(chatStreamSegments) && chatStreamSegments.length > 0) || + !!String(chatStream || "").trim(); + const fbLine = chatStream || extractWsAssistantText(payload.message || {}); + const fbTrim = String(fbLine || "").trim(); + const fbOk = !!fbTrim && !_isSilentReplyStream(fbTrim); + const n = await loadMessagesForActive({ keepDomIfNoHistoryRows: hadStream }); + const needWsFallback = + hadStream && fbOk && (n < 0 || !!loadMessagesForActive._needsWsTextFallback); + if (n >= 0) { + try { + if (streamRow && streamRow.parentNode) streamRow.remove(); + } catch (_) {} + } + if (needWsFallback) { + if (n < 0) { + try { + if (streamRow && streamRow.parentNode) streamRow.remove(); + } catch (_) {} + } + ok = await appendFinalAssistant(payload.message, fbLine); + } else if (n > 0) { + ok = true; + } else { + ok = n === 0; + } scrollMessagesToBottom(true); - ok = true; - } else if (sawStreamToolRefAttachments) { - // Streaming UI cannot render image_ref/relay_pointer; recover from persisted history so images appear - // without requiring a manual refresh. - try { - if (streamRow && streamRow.parentNode) streamRow.remove(); - } catch (_) {} - await loadMessagesForActive(); - scrollMessagesToBottom(true); - ok = true; } else { ok = await appendFinalAssistant(payload.message, chatStream || extractWsAssistantText(payload.message || {})); } diff --git a/interfaces/channels/wecom/longconn_runner.py b/interfaces/channels/wecom/longconn_runner.py index 78e1f274..27e94dfe 100644 --- a/interfaces/channels/wecom/longconn_runner.py +++ b/interfaces/channels/wecom/longconn_runner.py @@ -16,6 +16,7 @@ from interfaces.channels.wecom.normalize import normalize_wecom_event, normalize from svc.config.paths import db_path from svc.integrations.wecom_client import WeComClient from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def _safe_json(obj: Any) -> str: @@ -649,7 +650,7 @@ def _run_ws_forever(*, sender: WeComClient, deliver_outbound: bool, use_response def run_forever() -> int: - store = SqliteStore(db_path()) + store = get_assistant_store() sender = WeComClient(store) lock = _SingleInstanceLock(Path(db_path()).resolve().parent / "locks" / "wecom_longconn.lock") lock.acquire() diff --git a/interfaces/gateway/http_adapter.py b/interfaces/gateway/http_adapter.py index d0423376..4d160952 100644 --- a/interfaces/gateway/http_adapter.py +++ b/interfaces/gateway/http_adapter.py @@ -2,15 +2,14 @@ from __future__ import annotations from typing import Any -from svc.config.paths import db_path -from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from .context_builder import build_common_gateway_context from .dispatcher import build_gateway_method_handlers def _build_http_context() -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() context = build_common_gateway_context(store=store) context["session_event_subscribers"] = set() context["session_message_subscribers"] = {} diff --git a/interfaces/http/fastapi_app.py b/interfaces/http/fastapi_app.py index ba088f06..202d2501 100644 --- a/interfaces/http/fastapi_app.py +++ b/interfaces/http/fastapi_app.py @@ -32,7 +32,7 @@ from runtime.hooks_runtime import ( from runtime.skills import skill_runtime_diagnostics from runtime.prompt_prebuild import run_runtime_prewarm from svc.config.paths import PROJECT_ROOT, db_path -from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from interfaces.http.weixin_ilink_api import router as weixin_ilink_router from runtime.workspaces.experts import warm_expert_workspace_cache @@ -162,7 +162,7 @@ def _run_startup_hooks(app: FastAPI) -> None: # Force re-login after every gateway restart. try: - revoked = SqliteStore(db_path()).revoke_all_auth_sessions() + revoked = get_assistant_store().revoke_all_auth_sessions() if revoked > 0: _log_info(f"[auth] revoked sessions on startup: {revoked}") except Exception as exc: diff --git a/interfaces/ws/auth_and_hello.py b/interfaces/ws/auth_and_hello.py index 1c76dabb..bcbe46b1 100644 --- a/interfaces/ws/auth_and_hello.py +++ b/interfaces/ws/auth_and_hello.py @@ -6,6 +6,7 @@ from typing import Any from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def build_hello_ok_payload( @@ -50,7 +51,7 @@ def resolve_ws_auth(connect_params: dict[str, Any] | None) -> dict[str, Any]: token = device_token or bootstrap_token if not token: return {} - store = SqliteStore(db_path()) + store = get_assistant_store() token_hash = hashlib.sha256(token.encode("utf-8", errors="ignore")).hexdigest() session = store.get_auth_session(session_token_hash=token_hash) if not session or session.get("revoked_at"): diff --git a/interfaces/ws/server_methods_bridge.py b/interfaces/ws/server_methods_bridge.py index 776e64be..1e8959b1 100644 --- a/interfaces/ws/server_methods_bridge.py +++ b/interfaces/ws/server_methods_bridge.py @@ -6,6 +6,7 @@ from typing import Any, Callable from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from interfaces.gateway.context_builder import build_common_gateway_context @@ -22,7 +23,7 @@ def build_gateway_context( validate_relay_share_envelope: Callable[[dict[str, Any]], tuple[bool, str, dict[str, Any]]], now_ms: Callable[[], int], ) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() def _abort_chat_run(run_id: str) -> bool: rid = str(run_id or "").strip() diff --git a/interfaces/ws/turn_runner.py b/interfaces/ws/turn_runner.py index 8aa704da..08d9ad25 100644 --- a/interfaces/ws/turn_runner.py +++ b/interfaces/ws/turn_runner.py @@ -2,16 +2,21 @@ from __future__ import annotations import asyncio import json +import logging import threading import uuid from datetime import datetime from typing import Any, Callable from runtime.agents.factory import build_gateway_executor +from runtime.chat.persist_terminal_fallback import persist_assistant_text_if_turn_missing from runtime.gateway import OclawGateway from runtime.types import StandardMessage from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store + +_LOG = logging.getLogger(__name__) def _persisted_chat_attachments_nonempty(raw: Any) -> bool: @@ -36,6 +41,41 @@ def _persisted_chat_attachments_nonempty(raw: Any) -> bool: return False +def _assistant_row_meaningful_for_terminal_snapshot(m: Any) -> bool: + """Whether an assistant row should participate in WS ``final`` snapshot selection. + + Rows with ``content`` empty but ``tool_calls`` + ``event_payload.reasoning_content`` (common right + before tool execution finishes) were previously skipped, so the bridge picked an older assistant + row or only token fallback — matching \"reasoning then nothing\" when the post-tool reply never + landed in DB in time. + """ + if str(getattr(m, "role", "") or "").lower() != "assistant": + return False + if str(getattr(m, "content", "") or "").strip(): + return True + if _persisted_chat_attachments_nonempty(getattr(m, "attachments", None)): + return True + tc = getattr(m, "tool_calls", None) + if tc is not None: + s = str(tc).strip() + if not s or s.lower() == "null": + pass + else: + try: + parsed = json.loads(s) if isinstance(tc, str) else tc + except Exception: + return True + if isinstance(parsed, list) and len(parsed) > 0: + return True + if isinstance(parsed, dict) and parsed: + return True + ep = getattr(m, "event_payload", None) + if ep is None: + return False + es = str(ep).strip() + return bool(es) and es.lower() != "null" + + async def run_agent_turn_via_bridge( *, conn: Any, @@ -118,7 +158,7 @@ async def run_agent_turn_via_bridge( plan_agent_version = str(p.get("plan_agent_version") or "v1").strip().lower() or "v1" if plan_agent_version not in {"v1", "v2"}: plan_agent_version = "v1" - store = SqliteStore(db_path()) + store = get_assistant_store() gw = OclawGateway(store=store) ctx = conn.auth_ctx or {} @@ -341,11 +381,15 @@ async def run_agent_turn_via_bridge( session_id=str(session_id), role="assistant", content=str(user_facing_error), - turn_uuid=str(run_id_holder.get("run_id") or "") or None, + turn_uuid=None, event_type="assistant_text", ) except Exception: - pass + _LOG.exception( + "turn_runner_failed_reply_persist_failed session_id=%s run_id=%s", + str(session_id), + str(run_id_holder.get("run_id") or ""), + ) if send_response: await conn.send_res( req_id, @@ -531,8 +575,10 @@ async def run_agent_turn_via_bridge( if not final_text: with buf_lock: final_text = "".join(token_chunks) + stream_snapshot = str(final_text or "").strip() + turn_for_persist = str(getattr(result, "turn_uuid", "") or "").strip() final_msg: dict[str, Any] = {"role": "assistant", "content": final_text, "timestamp": now_ms()} - if run_status == "failed" and not str(final_text or "").strip(): + if run_status == "failed" and not stream_snapshot: err_code = run_last_error_code or "unknown_error" stop_reason = run_stop_reason or "failed" detail_line = f"\n详细原因:{run_error_detail}" if run_error_detail else "" @@ -549,11 +595,15 @@ async def run_agent_turn_via_bridge( session_id=str(session_id), role="assistant", content=str(final_text), - turn_uuid=str(rid or "") or None, + turn_uuid=turn_for_persist or None, event_type="assistant_text", ) except Exception: - pass + _LOG.exception( + "turn_runner_failed_final_reply_persist_failed session_id=%s run_id=%s", + str(session_id), + str(rid or ""), + ) elif run_status != "failed": try: def _event_payload_as_dict(raw: Any) -> dict[str, Any]: @@ -571,11 +621,12 @@ async def run_agent_turn_via_bridge( for m in reversed(list(persisted or [])): if str(getattr(m, "role", "") or "").lower() != "assistant": continue + if not _assistant_row_meaningful_for_terminal_snapshot(m): + continue content = str(getattr(m, "content", "") or "") atraw = getattr(m, "attachments", None) - if not content.strip() and not _persisted_chat_attachments_nonempty(atraw): - continue - final_text = content + # Keep streamed / gateway reply_text when newest DB row is tool-only (empty body). + final_text = str(content or "").strip() or str(final_text or "").strip() et = str(getattr(m, "event_type", "") or "").strip() ep = getattr(m, "event_payload", None) final_msg = { @@ -615,6 +666,14 @@ async def run_agent_turn_via_bridge( except Exception: pass + persist_assistant_text_if_turn_missing( + store=store, + session_id=str(session_id), + turn_uuid=turn_for_persist, + final_text=stream_snapshot, + log_prefix="turn_runner_ws_text_fallback_persisted", + ) + await conn.emit_chat_event(run_id=rid, state="final", reply=str(final_text or ""), message=final_msg, session_key=str(session_id)) try: await conn.send_event( diff --git a/requirements.txt b/requirements.txt index cb708b13..24cb7787 100644 --- a/requirements.txt +++ b/requirements.txt @@ -18,3 +18,6 @@ cryptography>=42.0.0 anthropic PyYAML>=6.0.0 python-dotenv>=1.0.0 +sqlalchemy>=2.0.0 +psycopg[binary]>=3.1.0 +alembic>=1.13.0 diff --git a/runtime/agent_core_attempt.py b/runtime/agent_core_attempt.py index 088126e3..06dc1e80 100644 --- a/runtime/agent_core_attempt.py +++ b/runtime/agent_core_attempt.py @@ -224,7 +224,7 @@ def run_attempt(*, store: Any, data: AttemptRunnerInput) -> AttemptRunnerOutput: final_text="", tool_traces=tuple(), handoff_note=f"{err_code}:{reason}", - turn_uuid="", + turn_uuid=str(data.turn_uuid or "").strip(), ), ) diff --git a/runtime/application/gateway/inbound_service.py b/runtime/application/gateway/inbound_service.py index 844bc7a4..7d6b12a4 100644 --- a/runtime/application/gateway/inbound_service.py +++ b/runtime/application/gateway/inbound_service.py @@ -7,6 +7,7 @@ from typing import Any from interfaces.channels.base import InboundMessage, OutboundMessage from interfaces.channels.wecom.wecom_bridge import WeComAdapter from runtime.types import normalize_interaction_mode, normalize_requested_specialist +from svc.persistence.assistant_store import get_assistant_store _CHANNEL_DISPATCH_INTERACTION_KEY_PREFIX = "channel.dispatch.interaction_mode." _CHANNEL_DISPATCH_SPECIALIST_KEY_PREFIX = "channel.dispatch.specialist." @@ -74,7 +75,7 @@ def _handle_productivity_commands(*, text: str, tenant_id: str, user_id: str) -> from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore - store = SqliteStore(db_path()) + store = get_assistant_store() if t.startswith("记待办 "): title = t[len("记待办 ") :].strip() @@ -508,7 +509,7 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]: from svc.persistence.sqlite_store import SqliteStore from svc.config.paths import db_path - store = SqliteStore(db_path()) + store = get_assistant_store() if channel_name == "wecom": account_id = _resolve_wecom_account_id(inbound, payload) or str(store.get_setting("wecom_bot_id") or "").strip() else: diff --git a/runtime/chat/persist_terminal_fallback.py b/runtime/chat/persist_terminal_fallback.py new file mode 100644 index 00000000..cda90f24 --- /dev/null +++ b/runtime/chat/persist_terminal_fallback.py @@ -0,0 +1,46 @@ +"""When the model run produced visible text but no ``chat_message`` row for this ``turn_uuid``.""" + +from __future__ import annotations + +import logging +from typing import Any + +_LOG = logging.getLogger(__name__) + + +def persist_assistant_text_if_turn_missing( + *, + store: Any, + session_id: str, + turn_uuid: str, + final_text: str, + log_prefix: str, +) -> bool: + """Insert one ``assistant_text`` row if none exists for ``turn_uuid``. Return True if inserted.""" + tu = str(turn_uuid or "").strip() + body = str(final_text or "").strip() + if not tu or not body: + return False + try: + rows = store.get_messages(session_id=str(session_id), limit=400) + if any( + str(getattr(m, "role", "") or "").lower() == "assistant" + and str(getattr(m, "turn_uuid", "") or "").strip() == tu + for m in (rows or []) + ): + return False + store.add_message( + session_id=str(session_id), + role="assistant", + content=body, + turn_uuid=tu, + event_type="assistant_text", + ) + _LOG.warning("%s session_id=%s turn_uuid=%s chars=%d", log_prefix, str(session_id), tu, len(body)) + return True + except Exception: + _LOG.exception("%s_failed session_id=%s turn_uuid=%s", log_prefix, str(session_id), tu) + return False + + +__all__ = ["persist_assistant_text_if_turn_missing"] diff --git a/runtime/gateway.py b/runtime/gateway.py index efad3e7e..6ce3453a 100644 --- a/runtime/gateway.py +++ b/runtime/gateway.py @@ -90,6 +90,8 @@ class OclawGatewayResult: relay_ttl_turn_count: int = 0 relay_ttl_session_count: int = 0 relay_ttl_keep_count: int = 0 + # agent-core 本轮 ``chat_message.turn_uuid``;供 WS 收尾与落库兜底对齐 + turn_uuid: str = "" @dataclass(frozen=True) @@ -889,15 +891,16 @@ class OclawGateway: ) self.store.add_message( session_id=msg.session_id, - tenant_id=msg.tenant_id, - user_id=msg.user_id, role="assistant", content=assignment_text, tool_calls=None, event_type="reasoning", ) except Exception: - pass + logger.exception( + "comprehensive_mode_assignment_message_persist_failed session_id=%s", + str(getattr(msg, "session_id", "") or ""), + ) # Build specialist/dynamic executor and dispatch only manager instruction to it. specialist_input_msg = StandardMessage( session_id=msg.session_id, @@ -1005,6 +1008,7 @@ class OclawGateway: relay_ttl_turn_count=int(ttl_stats.get("turn") or 0), relay_ttl_session_count=int(ttl_stats.get("session") or 0), relay_ttl_keep_count=int(ttl_stats.get("keep") or 0), + turn_uuid="", ) if action == "run_agent": system_prompt_override = str(shadow.decision.system_prompt_override or "") @@ -1146,7 +1150,9 @@ class OclawGateway: relay_ttl_turn_count=int(ttl_stats.get("turn") or 0), relay_ttl_session_count=int(ttl_stats.get("session") or 0), relay_ttl_keep_count=int(ttl_stats.get("keep") or 0), + turn_uuid="", ) + executed_turn_uuid = "" try: model = getattr(selected_executor, "model", None) tools = tools_override if tools_override is not None else getattr(selected_executor, "tools", None) @@ -1196,7 +1202,10 @@ class OclawGateway: max_tool_workers=_get_int_setting("AIA_TURN_MAX_TOOL_WORKERS", 8, 1, 32), max_attempts=_get_int_setting("AIA_OCLAW_MAX_ATTEMPTS", 2, 1, 5), memory_context=memory_context, - on_token=(None if interaction_mode == "comprehensive" else on_token), + # Always stream specialist tokens (incl. reasoning deltas) to WS clients. + # Comprehensive used to pass None here to hide raw specialist output before manager polish; + # that made admin webchat look like the stream died mid-"推理" until final/manager only. + on_token=on_token, on_progress=on_progress, on_tool_ui=on_tool_ui, should_stop=should_stop, @@ -1214,7 +1223,7 @@ class OclawGateway: specialist=manager_specialist, specialist_reply=specialist_reply, memory_enabled=memory_enabled, - on_token=on_token, + on_token=None, ) if self._looks_like_manager_instruction(reply, manager_instruction_text): if not self._looks_like_manager_instruction(specialist_reply, manager_instruction_text): @@ -1270,6 +1279,7 @@ class OclawGateway: relay_ttl_turn_count=int(ttl_stats.get("turn") or 0), relay_ttl_session_count=int(ttl_stats.get("session") or 0), relay_ttl_keep_count=int(ttl_stats.get("keep") or 0), + turn_uuid=str(executed_turn_uuid or ""), ) diff --git a/runtime/hooks/bundled/session-memory/handler.py b/runtime/hooks/bundled/session-memory/handler.py index 4168058d..231df248 100644 --- a/runtime/hooks/bundled/session-memory/handler.py +++ b/runtime/hooks/bundled/session-memory/handler.py @@ -7,6 +7,8 @@ import sys from pathlib import Path from typing import Any, Dict, List, Optional, Tuple +from svc.persistence.assistant_store import get_assistant_store + HOOK_KEY = "session-memory" @@ -125,7 +127,7 @@ def handle(event: Any) -> None: from svc.config.paths import db_path # type: ignore from svc.persistence.sqlite_store import SqliteStore # type: ignore - store = SqliteStore(db_path()) + store = get_assistant_store() msgs = store.get_messages(session_id=session_id, limit=max_msgs_n) except Exception: msgs = [] diff --git a/runtime/operations/providers/wecom.py b/runtime/operations/providers/wecom.py index cde3b2b6..82ad2402 100644 --- a/runtime/operations/providers/wecom.py +++ b/runtime/operations/providers/wecom.py @@ -9,6 +9,7 @@ from argparse import _SubParsersAction from interfaces.channels.wecom.longconn_runner import run_forever as run_wecom_longconn from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from .base import ChannelProvider @@ -32,7 +33,7 @@ _CLEAR_KEYS = [ def _store() -> SqliteStore: - return SqliteStore(db_path()) + return get_assistant_store() class WecomProvider(ChannelProvider): diff --git a/runtime/operations/runtime.py b/runtime/operations/runtime.py index 751a6320..f3b7622f 100644 --- a/runtime/operations/runtime.py +++ b/runtime/operations/runtime.py @@ -23,6 +23,15 @@ def _runtime_log_dir() -> Path: return Path(db_path()).resolve().parent / "logs" +def assistant_runtime_log_dir() -> Path: + """Directory for gateway/channel logs written by :func:`start_service`. + + Same rule as ``AIA_RUNTIME_LOG_DIR`` or ``)/logs``. + Use this from shell scripts so paths match Python ``stack up`` / ``start_service``. + """ + return _runtime_log_dir() + + def _read_state() -> dict[str, Any]: p = _runtime_file() if not p.exists(): diff --git a/runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh b/runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh new file mode 100644 index 00000000..03890f0d --- /dev/null +++ b/runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh @@ -0,0 +1,20 @@ +#!/usr/bin/env bash +# Import assistant SQLite (db_path) into PostgreSQL. Prerequisite: alembic upgrade head on target PG. +# +# Reads _local/system.env when present (--load-system-env). Target URL: AIA_ASSISTANT_DATABASE_URL +# (or pass extra args, e.g. --pg-url 'postgresql+psycopg://...' --dry-run). +# +# Usage: +# chmod +x runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh +# export AIA_ASSISTANT_DATABASE_URL='postgresql+psycopg://user:pass@host:5432/oclaw' +# ./runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh --dry-run +# ./runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh + +set -euo pipefail +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)" +cd "$ROOT" +export PYTHONPATH="${ROOT}${PYTHONPATH:+:${PYTHONPATH}}" +exec python "${ROOT}/runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py" \ + --load-system-env \ + --sqlite-from-db-path \ + "$@" diff --git a/runtime/operations/scripts/clear_all_chat_sessions.py b/runtime/operations/scripts/clear_all_chat_sessions.py new file mode 100644 index 00000000..f8bf3d88 --- /dev/null +++ b/runtime/operations/scripts/clear_all_chat_sessions.py @@ -0,0 +1,147 @@ +"""Delete every row in ``chat_session`` (and session-bound helper rows). + +Uses the same assistant store as the gateway (SQLite or PostgreSQL per env). +Requires ``--yes`` **and** environment ``AIA_CONFIRM_CHAT_SESSION_WIPE=1`` to avoid accidental wipes. + +PostgreSQL (force, avoids wiping SQLite by mistake):: + + set AIA_CONFIRM_CHAT_SESSION_WIPE=1 + python runtime/operations/scripts/clear_all_chat_sessions.py --yes --postgresql + +Or use ``runtime/operations/scripts/clear_postgres_chat_sessions.ps1`` (loads ``_local/system.env`` then runs the above). +""" + +from __future__ import annotations + +import argparse +import os +import sys +from typing import Any + + +def _exec(store: Any, sql: str) -> None: + with store._connect() as conn: + conn.execute(sql) + + +def _try_exec(store: Any, sql: str) -> bool: + try: + _exec(store, sql) + return True + except Exception: + return False + + +def main() -> int: + p = argparse.ArgumentParser(description=__doc__) + p.add_argument( + "--yes", + action="store_true", + help="Confirm destructive delete of all chat sessions.", + ) + p.add_argument( + "--dry-run", + action="store_true", + help="Only print how many sessions exist; do not delete.", + ) + p.add_argument( + "--postgresql", + action="store_true", + help="After loading env, force AIA_ASSISTANT_DB_BACKEND=postgresql and abort unless the store is PG.", + ) + args = p.parse_args() + if not args.yes and not args.dry_run: + print("Refusing to run without --yes (or use --dry-run to count only).", file=sys.stderr) + return 2 + if args.yes and not args.dry_run and str(os.getenv("AIA_CONFIRM_CHAT_SESSION_WIPE") or "").strip().lower() not in { + "1", + "true", + "yes", + "on", + }: + print( + "Refusing destructive wipe: set environment AIA_CONFIRM_CHAT_SESSION_WIPE=1 together with --yes.", + file=sys.stderr, + ) + return 2 + + try: + from interfaces.http.fastapi_app import load_system_env + + load_system_env() + except Exception: + pass + + if args.postgresql: + os.environ["AIA_ASSISTANT_DB_BACKEND"] = "postgresql" + + from svc.persistence.assistant_store import get_assistant_store, reset_assistant_store_singleton + + if args.postgresql: + reset_assistant_store_singleton() + + store = get_assistant_store() + if args.postgresql and not bool(getattr(store, "_use_pg", False)): + print( + "error: --postgresql was set but assistant store is not PostgreSQL " + "(check AIA_ASSISTANT_DATABASE_URL / OPS_ASSISTANT_DATABASE_URL).", + file=sys.stderr, + ) + return 2 + n0 = int(store.count_sessions() or 0) + print(f"session_count_before={n0}") + if args.dry_run: + return 0 + if n0 <= 0: + print("nothing_to_do") + return 0 + + # Rows that reference sessions but are not always ON DELETE CASCADE across backends. + for sql in ( + "DELETE FROM trace_event", + "DELETE FROM agent_eval_log", + "DELETE FROM oclaw_attempt", + "DELETE FROM oclaw_run", + "DELETE FROM oclaw_task", + "DELETE FROM memory_vector WHERE memory_id IN (SELECT memory_id FROM memory_item WHERE session_id IN (SELECT id FROM chat_session))", + "DELETE FROM memory_item WHERE session_id IN (SELECT id FROM chat_session)", + "DELETE FROM memory_hit_log WHERE session_id IN (SELECT id FROM chat_session)", + ): + if _try_exec(store, sql): + print(f"ok_stmt={sql[:72]}...") + else: + print(f"skip_stmt={sql[:72]}...") + + deleted_bulk = 0 + try: + with store._connect() as conn: + cur = conn.execute("DELETE FROM chat_session") + deleted_bulk = int(getattr(cur, "rowcount", 0) or 0) + except Exception as exc: + print(f"bulk_delete_chat_session_failed={exc!r}; falling back to per-session delete") + batch = 0 + while True: + rows = store.list_sessions(limit=400, offset=0) + if not rows: + break + for s in rows: + store.delete_session(str(s.id)) + batch += 1 + if batch > 1_000_000: + print("abort_loop_guard", file=sys.stderr) + return 1 + deleted_bulk = batch + + n1 = int(store.count_sessions() or 0) + print(f"deleted_sessions_bulk={deleted_bulk}") + print(f"session_count_after={n1}") + try: + store._chat_messages_repo().delete_messages_where_session_missing() + store._tool_log_queries_repo().delete_tool_logs_where_session_missing() + except Exception as exc: + print(f"orphan_cleanup_note={exc!r}") + return 0 if n1 == 0 else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/runtime/operations/scripts/clear_postgres_chat_sessions.ps1 b/runtime/operations/scripts/clear_postgres_chat_sessions.ps1 new file mode 100644 index 00000000..e970db80 --- /dev/null +++ b/runtime/operations/scripts/clear_postgres_chat_sessions.ps1 @@ -0,0 +1,71 @@ +# Wipe all chat_session rows (and related session-bound rows) on PostgreSQL only. +# Loads _local/system.env into the process, forces PG backend, sets confirmation env. +# Usage (from repo root is fine): +# .\runtime\operations\scripts\clear_postgres_chat_sessions.ps1 +# Dry-run (count only): +# .\runtime\operations\scripts\clear_postgres_chat_sessions.ps1 -DryRun + +param( + [switch]$DryRun = $false +) + +$ErrorActionPreference = "Stop" + +function Import-DotEnvFile([string]$path) { + if (-not (Test-Path $path)) { return } + Get-Content -LiteralPath $path -Encoding UTF8 | ForEach-Object { + $line = $_.Trim() + if (-not $line -or $line.StartsWith("#")) { return } + $idx = $line.IndexOf("=") + if ($idx -lt 1) { return } + $k = $line.Substring(0, $idx).Trim() + $v = $line.Substring($idx + 1).Trim() + if ($k) { + [System.Environment]::SetEnvironmentVariable($k, $v, "Process") + } + } +} + +$repoRoot = $null +$cur = (Resolve-Path $PSScriptRoot).Path +for ($i = 0; $i -lt 16; $i++) { + if (Test-Path (Join-Path $cur "oclaw.json")) { + $repoRoot = $cur + break + } + $parent = Split-Path -Parent $cur + if (-not $parent -or $parent -eq $cur) { break } + $cur = $parent +} +if (-not $repoRoot) { + throw "Could not find oclaw.json above $PSScriptRoot" +} +Set-Location $repoRoot +$env:PYTHONPATH = $repoRoot + +$envFile = Join-Path $repoRoot "_local\system.env" +Import-DotEnvFile $envFile + +$env:AIA_ASSISTANT_DB_BACKEND = "postgresql" +$env:AIA_CONFIRM_CHAT_SESSION_WIPE = "1" + +$venvPython = Join-Path $repoRoot ".venv\Scripts\python.exe" +$pythonExe = $(if (Test-Path $venvPython) { $venvPython } else { "python" }) + +$args = @( + (Join-Path $repoRoot "runtime\operations\scripts\clear_all_chat_sessions.py") +) +if ($DryRun) { + $args += "--dry-run" +} else { + $args += "--yes" +} +$args += "--postgresql" + +Write-Host "repo=$repoRoot" -ForegroundColor Cyan +Write-Host "env_file=$envFile" -ForegroundColor DarkGray +Write-Host "python=$pythonExe" -ForegroundColor DarkGray +Write-Host "dry_run=$DryRun" -ForegroundColor DarkGray + +& $pythonExe @args +exit $LASTEXITCODE diff --git a/runtime/operations/scripts/create_bind_code.py b/runtime/operations/scripts/create_bind_code.py index 6231f55a..e910bdce 100644 --- a/runtime/operations/scripts/create_bind_code.py +++ b/runtime/operations/scripts/create_bind_code.py @@ -5,10 +5,11 @@ import sys from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def main() -> int: - store = SqliteStore(db_path()) + store = get_assistant_store() tenants = store.list_tenants(limit=1) if tenants: tenant_id = tenants[0]["id"] diff --git a/runtime/operations/scripts/cutover_sqlite_to_postgresql.ps1 b/runtime/operations/scripts/cutover_sqlite_to_postgresql.ps1 new file mode 100644 index 00000000..343a14d3 --- /dev/null +++ b/runtime/operations/scripts/cutover_sqlite_to_postgresql.ps1 @@ -0,0 +1,128 @@ +<# +.SYNOPSIS + Backup assistant SQLite DB, optionally dry-run, then import into PostgreSQL (empty schema). + +.DESCRIPTION + 1) Resolves SQLite path (explicit -SqlitePath or via Python db_path() / AIA_ASSISTANT_DB_PATH). + 2) Copies the file to data/pg_cutover_backups/ (or -BackupDir). + 3) Runs migrate_assistant_sqlite_to_postgresql.py --dry-run unless -SkipDryRun. + 4) Runs the same script without --dry-run unless -DryRunOnly. + + Set -PgUrl and/or ensure AIA_ASSISTANT_DATABASE_URL is set; with -LoadSystemEnv, URL may come only from _local/system.env (then -PgUrl can be omitted). + +.EXAMPLE + .\cutover_sqlite_to_postgresql.ps1 -PgUrl "postgresql+psycopg://postgres:PASS@127.0.0.1:5432/oclaw" + +.EXAMPLE + $env:AIA_ASSISTANT_DATABASE_URL = "postgresql+psycopg://..." + .\cutover_sqlite_to_postgresql.ps1 + +.EXAMPLE + .\cutover_sqlite_to_postgresql.ps1 -SqlitePath "D:\data\ai_ops.sqlite" -PgUrl "postgresql://..." -DryRunOnly +#> +[CmdletBinding()] +param( + [string] $PgUrl = "", + [string] $SqlitePath = "", + [string] $BackupDir = "", + [switch] $LoadSystemEnv, + [switch] $SkipDryRun, + [switch] $DryRunOnly, + [switch] $NoBackup +) + +$ErrorActionPreference = "Stop" +$RepoRoot = (Resolve-Path (Join-Path $PSScriptRoot "..\..\..")).Path +Set-Location $RepoRoot + +if ($LoadSystemEnv) { + $envFile = Join-Path $RepoRoot "_local\system.env" + if (Test-Path -LiteralPath $envFile) { + Get-Content -LiteralPath $envFile | ForEach-Object { + $line = $_.Trim() + if (-not $line -or $line.StartsWith("#")) { return } + $i = $line.IndexOf("=") + if ($i -lt 1) { return } + $k = $line.Substring(0, $i).Trim() + $v = $line.Substring($i + 1).Trim() + if ($k) { Set-Item -Path "Env:$k" -Value $v } + } + Write-Host "Loaded _local/system.env into process environment." + } + else { + Write-Warning "LoadSystemEnv specified but $envFile not found." + } +} + +if (-not $PgUrl) { + $PgUrl = $env:AIA_ASSISTANT_DATABASE_URL +} +if (-not $PgUrl -and -not $LoadSystemEnv) { + throw "Provide -PgUrl, set AIA_ASSISTANT_DATABASE_URL, or use -LoadSystemEnv (PostgreSQL URL in _local/system.env)." +} + +if (-not $SqlitePath) { + $env:_OC_REPO_ROOT_FOR_PY = $RepoRoot + try { + $SqlitePath = (& python -c "import os,sys; sys.path.insert(0, os.environ['_OC_REPO_ROOT_FOR_PY']); from svc.config.paths import db_path; print(db_path(), end='')").Trim() + } + finally { + Remove-Item Env:_OC_REPO_ROOT_FOR_PY -ErrorAction SilentlyContinue + } + if (-not $SqlitePath) { throw "Could not resolve SQLite path via db_path()." } + Write-Host "Resolved SQLite: $SqlitePath" +} + +if (-not (Test-Path -LiteralPath $SqlitePath)) { + throw "SQLite file not found: $SqlitePath" +} + +$resolvedSqlite = (Resolve-Path -LiteralPath $SqlitePath).Path +$migrate = Join-Path $RepoRoot "runtime\operations\scripts\migrate_assistant_sqlite_to_postgresql.py" +if (-not (Test-Path -LiteralPath $migrate)) { + throw "Migration script not found: $migrate" +} + +if (-not $BackupDir) { + $BackupDir = Join-Path $RepoRoot "data\pg_cutover_backups" +} +New-Item -ItemType Directory -Force -Path $BackupDir | Out-Null +$stamp = Get-Date -Format "yyyyMMdd_HHmmss" +$leaf = [System.IO.Path]::GetFileNameWithoutExtension($resolvedSqlite) +$bakName = "${leaf}_pre_pg_${stamp}.sqlite" +$bakPath = Join-Path $BackupDir $bakName + +if (-not $NoBackup) { + Copy-Item -LiteralPath $resolvedSqlite -Destination $bakPath -Force + Write-Host "Backup written: $bakPath" +} +else { + Write-Warning "NoBackup: skipping file copy (no SQLite backup created)." +} + +$common = @($migrate) +if ($LoadSystemEnv) { + $common += "--load-system-env" +} +$common += "--sqlite", $resolvedSqlite +if ($PgUrl) { + $common += @("--pg-url", $PgUrl) +} + +if (-not $SkipDryRun) { + Write-Host "=== Dry-run (row counts, no PG writes) ===" -ForegroundColor Cyan + & python @common "--dry-run" + if ($LASTEXITCODE -ne 0) { throw "Dry-run failed (exit $LASTEXITCODE)." } +} + +if ($DryRunOnly) { + Write-Host "DryRunOnly: skipping live import." -ForegroundColor Yellow + exit 0 +} + +Write-Host "=== Live import into PostgreSQL ===" -ForegroundColor Cyan +& python @common +if ($LASTEXITCODE -ne 0) { throw "Migration failed (exit $LASTEXITCODE)." } + +Write-Host "" +Write-Host "Done. Next: set AIA_ASSISTANT_DB_BACKEND=postgresql and AIA_ASSISTANT_DATABASE_URL in deployment, restart gateway." -ForegroundColor Green diff --git a/runtime/operations/scripts/cutover_sqlite_to_postgresql.sh b/runtime/operations/scripts/cutover_sqlite_to_postgresql.sh new file mode 100644 index 00000000..a5547595 --- /dev/null +++ b/runtime/operations/scripts/cutover_sqlite_to_postgresql.sh @@ -0,0 +1,67 @@ +#!/usr/bin/env bash +# Backup SQLite assistant DB, optional dry-run, then import into PostgreSQL (empty schema recommended). +# +# Loads _local/system.env via the Python migrator (--load-system-env) so AIA_ASSISTANT_DATABASE_URL / +# AIA_ASSISTANT_DB_PATH match the gateway. Override URL with extra args, e.g. --pg-url 'postgresql+...' +# +# Environment: +# SQLITE_PATH optional; default: db_path() after load_system_env +# BACKUP_DIR optional; default: /data/pg_cutover_backups +# NO_BACKUP=1 skip file copy +# SKIP_DRY_RUN=1 skip dry-run pass +# DRY_RUN_ONLY=1 only dry-run +# +# Example: +# chmod +x runtime/operations/scripts/cutover_sqlite_to_postgresql.sh +# ./runtime/operations/scripts/cutover_sqlite_to_postgresql.sh +# ./runtime/operations/scripts/cutover_sqlite_to_postgresql.sh --pg-url 'postgresql+psycopg://u:p@h:5432/oclaw' + +set -euo pipefail +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)" +cd "$ROOT" +export PYTHONPATH="${ROOT}${PYTHONPATH:+:${PYTHONPATH}}" +MIG="${ROOT}/runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py" + +resolve_sqlite() { + if [[ -n "${SQLITE_PATH:-}" ]]; then + printf '%s' "$SQLITE_PATH" + return + fi + python -c "import os,sys; sys.path.insert(0, r'''${ROOT}'''); os.chdir(r'''${ROOT}'''); from svc.config.bootstrap_env import load_system_env; load_system_env(force=True); from svc.config.paths import db_path; print(db_path(), end='')" +} + +SQLITE="$(resolve_sqlite)" +if [[ ! -f "$SQLITE" ]]; then + echo "SQLite file not found: $SQLITE" >&2 + exit 1 +fi + +BACKUP_DIR="${BACKUP_DIR:-${ROOT}/data/pg_cutover_backups}" +mkdir -p "$BACKUP_DIR" +STAMP="$(date +%Y%m%d_%H%M%S)" +BASE="$(basename "$SQLITE" .sqlite)" +BAK="${BACKUP_DIR}/${BASE}_pre_pg_${STAMP}.sqlite" + +if [[ "${NO_BACKUP:-}" != "1" ]]; then + cp -f "$SQLITE" "$BAK" + echo "Backup written: $BAK" +else + echo "NO_BACKUP=1: skipping SQLite file backup." >&2 +fi + +COMMON=(--load-system-env --sqlite "$SQLITE") + +if [[ "${SKIP_DRY_RUN:-}" != "1" ]]; then + echo "=== Dry-run (row counts, no PG writes) ===" >&2 + python "$MIG" "${COMMON[@]}" "$@" --dry-run +fi + +if [[ "${DRY_RUN_ONLY:-}" == "1" ]]; then + echo "DRY_RUN_ONLY=1: skipping live import." >&2 + exit 0 +fi + +echo "=== Live import into PostgreSQL ===" >&2 +python "$MIG" "${COMMON[@]}" "$@" +echo "" >&2 +echo "Done. Set AIA_ASSISTANT_DB_BACKEND=postgresql and AIA_ASSISTANT_DATABASE_URL, then restart the gateway." >&2 diff --git a/runtime/operations/scripts/install_extra_mcp_batch.py b/runtime/operations/scripts/install_extra_mcp_batch.py index 642d7e31..3f8e152e 100644 --- a/runtime/operations/scripts/install_extra_mcp_batch.py +++ b/runtime/operations/scripts/install_extra_mcp_batch.py @@ -13,6 +13,7 @@ sys.path.insert(0, str(ROOT)) from svc.config.paths import db_path # noqa: E402 from svc.persistence.sqlite_store import SqliteStore # noqa: E402 +from svc.persistence.assistant_store import get_assistant_store from runtime.tools.mcp.installer import McpServerManifest, _safe_server_id, install_mcp_server # noqa: E402 @@ -24,7 +25,7 @@ BATCH: list[tuple[str, str, list[str]]] = [ def main() -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() for seed, source_ref, entry_args in BATCH: server_id = _safe_server_id(seed) manifest = McpServerManifest( diff --git a/runtime/operations/scripts/install_mcp_context7.py b/runtime/operations/scripts/install_mcp_context7.py index 2827d939..1e539195 100644 --- a/runtime/operations/scripts/install_mcp_context7.py +++ b/runtime/operations/scripts/install_mcp_context7.py @@ -16,6 +16,7 @@ sys.path.insert(0, str(ROOT)) from svc.config.paths import db_path # noqa: E402 from svc.persistence.sqlite_store import SqliteStore # noqa: E402 +from svc.persistence.assistant_store import get_assistant_store from runtime.tools.mcp.installer import McpServerManifest, install_mcp_server # noqa: E402 from runtime.tools.mcp.runtime import McpProcessRuntime # noqa: E402 @@ -85,7 +86,7 @@ def _append_generalist_binding(store: SqliteStore, server_id: str) -> None: def main() -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() manifest = McpServerManifest( server_id="mcp-context7", source_type="npm", diff --git a/runtime/operations/scripts/install_mcp_sqlite_brave_calendar.py b/runtime/operations/scripts/install_mcp_sqlite_brave_calendar.py index f64df79a..125cf3a4 100644 --- a/runtime/operations/scripts/install_mcp_sqlite_brave_calendar.py +++ b/runtime/operations/scripts/install_mcp_sqlite_brave_calendar.py @@ -15,6 +15,7 @@ sys.path.insert(0, str(ROOT)) from svc.config.paths import db_path # noqa: E402 from svc.persistence.sqlite_store import SqliteStore # noqa: E402 +from svc.persistence.assistant_store import get_assistant_store from runtime.tools.mcp.installer import McpServerManifest, install_mcp_server # noqa: E402 from runtime.tools.mcp.runtime import McpProcessRuntime # noqa: E402 @@ -61,7 +62,7 @@ def _sync_tools(store: SqliteStore, server_id: str) -> bool: def main() -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() db_file = db_path() bundles: list[McpServerManifest] = [ diff --git a/runtime/operations/scripts/live_chat_probe.py b/runtime/operations/scripts/live_chat_probe.py new file mode 100644 index 00000000..898780ef --- /dev/null +++ b/runtime/operations/scripts/live_chat_probe.py @@ -0,0 +1,251 @@ +#!/usr/bin/env python3 +"""对「正在运行」的网关发一条与控制台 Chat 相同的 HTTP 消息,并打印模型 / 会话模式 / 响应 / 库尾行。 + +本机执行(需网关已启动,且与浏览器访问的是同一 ``--base-url``):: + + # 方式一:用密码登录(tenant 可省略,默认取库里第一个团队) + set AIA_LIVE_CHAT_PASSWORD=你的控制台密码 + python runtime/operations/scripts/live_chat_probe.py --session-id <浏览器地址栏里的会话id> + + # 方式二:从浏览器开发者工具复制 Bearer token + python runtime/operations/scripts/live_chat_probe.py --session-id --token + +可选:加 ``--dump-db`` 时,会再读 ``_local/system.env`` 里的 ``AIA_ASSISTANT_*``,直连同一 PG/SQLite 打印该会话最近几条 ``chat_message``(用于对照「HTTP 成功但库里没有 assistant」)。 + +注意:Cursor 里的 AI **不能**替你连本机浏览器或 WebSocket;此脚本等价于你自己在 Network 里手动 POST,只是自动带上当前「选用模型」与会话模式字段。 +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +from pathlib import Path +from typing import Any + + +def _repo_root() -> Path: + return Path(__file__).resolve().parents[3] + + +def _load_local_env_for_db_only() -> None: + root = _repo_root() + p = root / "_local" / "system.env" + if not p.is_file(): + return + for raw in p.read_text(encoding="utf-8", errors="replace").splitlines(): + s = raw.strip() + if not s or s.startswith("#") or "=" not in s: + continue + k, _, rest = s.partition("=") + key = k.strip() + if not key or key in os.environ: + continue + if not key.startswith("AIA_ASSISTANT_") and key not in ( + "OPS_ASSISTANT_DB_BACKEND", + "OPS_ASSISTANT_DATABASE_URL", + "OPS_ASSISTANT_DB_PATH", + ): + continue + val = rest.strip() + if len(val) >= 2 and val[0] == val[-1] and val[0] in "\"'": + val = val[1:-1] + os.environ[key] = val + + +def _dump_db_messages(session_id: str, limit: int) -> None: + _load_local_env_for_db_only() + root = _repo_root() + if str(root) not in sys.path: + sys.path.insert(0, str(root)) + + from sqlalchemy import func, select + + from svc.persistence.db.engine import clear_assistant_engine_cache, get_assistant_engine + from svc.persistence.assistant_store import reset_assistant_store_singleton + from svc.persistence.db.tables import chat_message + + clear_assistant_engine_cache() + reset_assistant_store_singleton() + eng = get_assistant_engine() + print("\n--- DB (same env as _local/system.env assistant store) ---") + print("engine:", str(eng.url).split("@")[-1]) + tc_len = func.length(func.coalesce(chat_message.c.tool_calls, "")) + stmt = ( + select( + chat_message.c.id, + chat_message.c.role, + chat_message.c.event_type, + func.length(func.coalesce(chat_message.c.content, "")).label("content_len"), + (tc_len > 2).label("has_tc"), + ) + .where(chat_message.c.session_id == session_id) + .order_by(chat_message.c.id.desc()) + .limit(limit) + ) + with eng.connect() as c: + n = c.execute( + select(func.count()).select_from(chat_message).where(chat_message.c.session_id == session_id) + ).scalar() + print("chat_message count:", int(n or 0)) + rows = list(c.execute(stmt).mappings()) + for r in reversed(rows): + print(dict(r)) + + +def _auth_header(token: str) -> dict[str, str]: + t = str(token or "").strip() + if t.lower().startswith("bearer "): + t = t[7:].strip() + return {"authorization": f"Bearer {t}"} + + +def main() -> int: + ap = argparse.ArgumentParser(description="Live HTTP chat probe against running gateway") + ap.add_argument("--base-url", default=os.getenv("AIA_LIVE_CHAT_BASE_URL", "http://127.0.0.1:8787").rstrip("/")) + ap.add_argument("--session-id", required=True, help="当前 Chat 会话 id(与浏览器一致)") + ap.add_argument("--text", default="请用一句话回复:当前探针在测试落库。", help="发送的正文") + ap.add_argument("--token", default=os.getenv("AIA_LIVE_CHAT_TOKEN", "").strip(), help="Bearer token(可设环境变量)") + ap.add_argument("--tenant-id", default=os.getenv("AIA_LIVE_CHAT_TENANT_ID", "").strip()) + ap.add_argument("--username", default=os.getenv("AIA_LIVE_CHAT_USERNAME", "administrator").strip()) + ap.add_argument("--password", default=os.getenv("AIA_LIVE_CHAT_PASSWORD", "").strip()) + ap.add_argument("--dump-db", action="store_true", help="发送后再读本地 assistant 库该会话消息尾") + ap.add_argument("--db-tail", type=int, default=8, help="--dump-db 时打印最近几条") + args = ap.parse_args() + + try: + import httpx + except ImportError: + print("需要 httpx: pip install httpx", file=sys.stderr) + return 2 + + sid = str(args.session_id).strip() + if not sid: + print("session-id 为空", file=sys.stderr) + return 2 + + base = str(args.base_url).strip().rstrip("/") + timeout = httpx.Timeout(300.0, connect=15.0) + + with httpx.Client(base_url=base, timeout=timeout) as client: + token = str(args.token or "").strip() + if not token: + pw = str(args.password or "").strip() + if not pw: + print( + "未提供 token:请设置 --password 或环境变量 AIA_LIVE_CHAT_PASSWORD," + "或 --token / AIA_LIVE_CHAT_TOKEN", + file=sys.stderr, + ) + return 2 + r0 = client.post("/admin/api/auth/bootstrap", json={}) + if r0.status_code != 200: + print("bootstrap", r0.status_code, r0.text, file=sys.stderr) + return 1 + body: dict[str, Any] = { + "username": str(args.username), + "password": pw, + "purpose": "console", + } + tid = str(args.tenant_id or "").strip() + if tid: + body["tenant_id"] = tid + lr = client.post("/admin/api/auth/login", json=body) + if lr.status_code != 200: + print("login http", lr.status_code, lr.text, file=sys.stderr) + return 1 + lj = lr.json() + if not lj.get("ok"): + print("login", lj, file=sys.stderr) + return 1 + token = str(lj.get("token") or "").strip() + if not token: + print("login 无 token", lj, file=sys.stderr) + return 1 + sess = lj.get("session") or {} + print( + "login ok tenant_id=", + str(sess.get("tenant_id") or ""), + "user_id=", + str(sess.get("user_id") or ""), + "username=", + str(sess.get("username") or ""), + ) + + h = _auth_header(token) + + mr = client.get("/admin/api/models", headers=h) + if mr.status_code != 200: + print("GET /models", mr.status_code, mr.text, file=sys.stderr) + return 1 + mj = mr.json() + if not mj.get("ok"): + print("models", mj, file=sys.stderr) + return 1 + active = str(mj.get("active_llm_profile_id") or "") + profiles = mj.get("profiles") or [] + name = "" + mode = "" + model = "" + for p in profiles: + if isinstance(p, dict) and str(p.get("id") or "") == active: + name = str(p.get("name") or "") + mode = str(p.get("mode") or "") + model = str(p.get("model") or "") + break + print("\n--- 当前选用模型(与控制台一致) ---") + print("active_llm_profile_id:", active) + print("profile:", name, "| mode:", mode, "| model:", model) + dbp = mj.get("db_path") + if dbp is not None: + print("gateway reports db_path:", dbp) + + gr = client.get(f"/admin/api/chat/sessions/{sid}/mode", headers=h) + if gr.status_code != 200: + print("GET session/mode", gr.status_code, gr.text, file=sys.stderr) + return 1 + gj = gr.json() + if not gj.get("ok"): + print("session/mode", gj, file=sys.stderr) + return 1 + print("\n--- 当前会话模式(将原样带入 POST) ---") + print(json.dumps({k: gj.get(k) for k in ("interaction_mode", "specialist", "memory_mode", "execution_mode", "confirm_strategy", "plan_agent_version") if k in gj}, ensure_ascii=False)) + + payload = { + "text": str(args.text), + "interaction_mode": gj.get("interaction_mode"), + "specialist": gj.get("specialist"), + "memory_mode": gj.get("memory_mode"), + "execution_mode": gj.get("execution_mode"), + } + pr = client.post(f"/admin/api/chat/sessions/{sid}/messages", headers=h, json=payload) + if pr.status_code != 200: + print("POST messages", pr.status_code, pr.text, file=sys.stderr) + return 1 + pj = pr.json() + print("\n--- POST /messages 响应 ---") + print(json.dumps(pj, ensure_ascii=False, indent=2)[:8000]) + if not pj.get("ok"): + return 1 + reply = str(pj.get("reply") or "") + print("\nreply 非空:", bool(reply.strip()), "len=", len(reply)) + + lr2 = client.get(f"/admin/api/chat/sessions/{sid}/messages?limit=50", headers=h) + if lr2.status_code == 200: + lj2 = lr2.json() + msgs = lj2.get("messages") or [] + roles = [str(m.get("role") or "") for m in msgs if isinstance(m, dict)] + print("\n--- GET messages(接口返回,最多 50 条) ---") + print("count=", len(msgs), "roles=", roles[-12:]) + + if args.dump_db: + try: + _dump_db_messages(sid, max(1, int(args.db_tail))) + except Exception as e: + print("\n--dump-db 失败(可忽略):", type(e).__name__, e, file=sys.stderr) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py b/runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py new file mode 100644 index 00000000..132c81f7 --- /dev/null +++ b/runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py @@ -0,0 +1,407 @@ +"""Copy assistant data from SQLite (db_path file) into PostgreSQL (schema from Alembic / bootstrap). + +**Prerequisite:** target PostgreSQL already has schema (``alembic upgrade head`` or +``svc/persistence/ddl/postgresql_bootstrap.sql``). This script copies **data only**. + +Tables are copied in **foreign-key safe order** (from SQLite ``PRAGMA foreign_key_list``), and only +**columns present in both** SQLite and PostgreSQL are inserted. + +By default the script aborts if any target table in ``public`` already has rows (empty PG only). +Use ``--allow-non-empty`` to skip that check (you are responsible for avoiding duplicates / FK errors). + +**PostgreSQL URL** is taken from ``--pg-url`` if set; otherwise from the first non-empty environment +variable among ``AIA_ASSISTANT_DATABASE_URL``, ``OPS_ASSISTANT_DATABASE_URL``, ``AIA_ASSISTANT_PG_DSN``, +``OPS_ASSISTANT_PG_DSN``. Use ``--load-system-env`` to merge ``_local/system.env`` first (same as the +HTTP gateway). + +**Open-source / headless device (Linux example)**:: + + export AIA_ASSISTANT_DATABASE_URL='postgresql+psycopg://USER:PASS@HOST:5432/oclaw' + export AIA_ASSISTANT_DB_PATH=/var/lib/oclaw/data/ai_ops.sqlite # optional; default data/ai_ops.sqlite under repo + cd /path/to/oclaw && PYTHONPATH=. python runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py \\ + --load-system-env --sqlite-from-db-path --dry-run + # then same without --dry-run + +Or use the wrapper script ``runtime/operations/scripts/assistant_import_sqlite_to_postgresql.sh``. + +**Explicit paths**:: + + python runtime/operations/scripts/migrate_assistant_sqlite_to_postgresql.py \\ + --sqlite data/ai_ops.sqlite \\ + --pg-url postgresql+psycopg://postgres:pass@127.0.0.1:5432/oclaw + +""" + +from __future__ import annotations + +import argparse +import os +import re +import sqlite3 +import sys +from collections import defaultdict +from pathlib import Path +from typing import Any, Iterable + +import psycopg +from psycopg.rows import dict_row + +_REPO_ROOT = Path(__file__).resolve().parents[3] +if str(_REPO_ROOT) not in sys.path: + sys.path.insert(0, str(_REPO_ROOT)) + +from svc.persistence.pg_adapter import normalize_psycopg_conninfo + + +def _pg_row_first_value(row: Any) -> Any: + if row is None: + raise ValueError("expected a row") + if isinstance(row, dict): + return next(iter(row.values())) + return row[0] + + +_IDENT = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*\Z") + + +def _require_ident(name: str) -> str: + if not _IDENT.fullmatch(name): + raise ValueError(f"invalid SQL identifier: {name!r}") + return name + + +def _sqlite_user_tables(sl: sqlite3.Connection) -> list[str]: + rows = sl.execute( + """ + SELECT name FROM sqlite_master + WHERE type='table' AND name NOT LIKE 'sqlite_%' + ORDER BY name + """ + ).fetchall() + return [str(r[0]) for r in rows] + + +def _pg_public_tables(pg: psycopg.Connection) -> set[str]: + with pg.cursor() as cur: + cur.execute( + """ + SELECT tablename FROM pg_catalog.pg_tables + WHERE schemaname = 'public' + """ + ) + return {str(_pg_row_first_value(r)) for r in cur.fetchall()} + + +def _fk_parents_for_table(sl: sqlite3.Connection, table: str) -> set[str]: + t = _require_ident(table) + rows = sl.execute(f"PRAGMA foreign_key_list({t})").fetchall() + out: set[str] = set() + for r in rows: + # (id, seq, table, from, to, on_update, on_delete, match) + ref = str(r[2]) + if ref: + out.add(ref) + return out + + +def _topological_sort(nodes: list[str], parents: dict[str, set[str]]) -> list[str]: + """``parents[t]`` = tables that must be copied *before* ``t`` (referenced by FK).""" + node_set = list(nodes) + seen = set(node_set) + if len(seen) != len(node_set): + raise ValueError("duplicate table in migration list") + + children: dict[str, list[str]] = defaultdict(list) + indegree: dict[str, int] = {} + for t in node_set: + ps = parents.get(t, set()) & seen + indegree[t] = len(ps) + for p in ps: + children[p].append(t) + for ch in children.values(): + ch.sort() + + queue = sorted([t for t in node_set if indegree[t] == 0]) + out: list[str] = [] + while queue: + n = queue.pop(0) + out.append(n) + for c in children[n]: + indegree[c] -= 1 + if indegree[c] == 0: + queue.append(c) + queue.sort() + if len(out) != len(seen): + remain = seen - set(out) + raise SystemExit( + "Cannot derive a foreign-key-safe copy order (cycle or unresolved FK). " + f"Remaining tables: {sorted(remain)}" + ) + return out + + +def _migration_order(sl: sqlite3.Connection, tables: list[str]) -> list[str]: + node_set = list(tables) + parents = {t: _fk_parents_for_table(sl, t) & set(node_set) for t in node_set} + return _topological_sort(node_set, parents) + + +def _sqlite_columns(sl: sqlite3.Connection, table: str) -> list[str]: + t = _require_ident(table) + rows = sl.execute(f"PRAGMA table_info({t})").fetchall() + # cid, name, type, notnull, dflt_value, pk + return [str(r[1]) for r in rows] + + +def _pg_columns(pg: psycopg.Connection, table: str) -> set[str]: + t = _require_ident(table) + with pg.cursor() as cur: + cur.execute( + """ + SELECT column_name FROM information_schema.columns + WHERE table_schema = 'public' AND table_name = %s + """, + (t,), + ) + return {str(_pg_row_first_value(r)) for r in cur.fetchall()} + + +def _common_columns(sl: sqlite3.Connection, pg: psycopg.Connection, table: str) -> list[str]: + sc = _sqlite_columns(sl, table) + pc = _pg_columns(pg, table) + return [c for c in sc if c in pc and _IDENT.fullmatch(c)] + + +def _assert_pg_tables_empty(pg: psycopg.Connection, tables: Iterable[str]) -> None: + with pg.cursor() as cur: + for t in sorted(set(tables)): + _require_ident(t) + cur.execute(f'SELECT COUNT(*) AS n FROM "{t}"') + row = cur.fetchone() + n = int(_pg_row_first_value(row)) + if n: + raise SystemExit( + f"Refusing to import: PostgreSQL table {t!r} already has {n} row(s). " + "Use an empty schema after alembic upgrade, or pass --allow-non-empty if you " + "really intend to append (duplicates / FK failures are your risk)." + ) + + +def _sqlite_row_counts(sl: sqlite3.Connection, tables: Iterable[str]) -> dict[str, int]: + out: dict[str, int] = {} + for t in tables: + _require_ident(t) + n = int(sl.execute(f"SELECT COUNT(*) FROM {_require_ident(t)}").fetchone()[0]) + out[t] = n + return out + + +def _pg_row_counts(pg: psycopg.Connection, tables: Iterable[str]) -> dict[str, int]: + out: dict[str, int] = {} + with pg.cursor() as cur: + for t in tables: + _require_ident(t) + cur.execute(f'SELECT COUNT(*) AS n FROM "{t}"') + row = cur.fetchone() + out[t] = int(_pg_row_first_value(row)) + return out + + +def _copy_table( + *, + sl: sqlite3.Connection, + pg: psycopg.Connection, + table: str, + cols: list[str], + dry_run: bool, + batch: int, +) -> int: + if not cols: + return 0 + t = _require_ident(table) + cur = sl.execute(f"SELECT {', '.join(_require_ident(c) for c in cols)} FROM {t}") + rows = cur.fetchall() + if not rows: + return 0 + if dry_run: + return len(rows) + col_sql = ", ".join(f'"{_require_ident(c)}"' for c in cols) + placeholders = ", ".join(["%s"] * len(cols)) + sql = f'INSERT INTO "{t}" ({col_sql}) VALUES ({placeholders})' + tuples = [tuple(r[c] for c in cols) for r in rows] + with pg.cursor() as pc: + for i in range(0, len(tuples), max(1, batch)): + chunk = tuples[i : i + max(1, batch)] + pc.executemany(sql, chunk) + return len(rows) + + +def _serial_columns(pg: psycopg.Connection) -> list[tuple[str, str]]: + """Tables/columns backed by a PostgreSQL sequence (BIGSERIAL etc.), for post-import setval.""" + with pg.cursor() as cur: + cur.execute( + """ + SELECT table_name, column_name + FROM information_schema.columns + WHERE table_schema = 'public' + AND column_default IS NOT NULL + AND column_default LIKE 'nextval%' + ORDER BY table_name, column_name + """ + ) + rows = cur.fetchall() + out: list[tuple[str, str]] = [] + for r in rows: + if isinstance(r, dict): + t = str(r["table_name"]) + c = str(r["column_name"]) + else: + t = str(r[0]) + c = str(r[1]) + if t and c: + out.append((_require_ident(t), _require_ident(c))) + return out + + +def _sync_sequences(pg: psycopg.Connection) -> None: + for t, col in _serial_columns(pg): + with pg.cursor() as cur: + cur.execute(f'SELECT COALESCE(MAX("{col}"), 1) AS mx FROM "{t}"') + row = cur.fetchone() + mx = int(_pg_row_first_value(row)) + try: + cur.execute( + "SELECT setval(pg_get_serial_sequence(%s, %s), %s, true)", + (t, col, mx), + ) + except Exception: + pass + + +def _pg_url_from_environ() -> str: + return ( + os.getenv("AIA_ASSISTANT_DATABASE_URL") + or os.getenv("OPS_ASSISTANT_DATABASE_URL") + or os.getenv("AIA_ASSISTANT_PG_DSN") + or os.getenv("OPS_ASSISTANT_PG_DSN") + or "" + ).strip() + + +def _resolve_sqlite_path(arg: str | None, use_db_path: bool) -> Path: + if use_db_path: + from svc.config.paths import db_path + + return Path(db_path()).expanduser().resolve() + if not arg: + raise SystemExit("Either pass --sqlite PATH or --sqlite-from-db-path") + return Path(arg).expanduser().resolve() + + +def main() -> None: + ap = argparse.ArgumentParser(description=__doc__) + ap.add_argument("--sqlite", default=None, help="Path to source SQLite assistant DB") + ap.add_argument( + "--sqlite-from-db-path", + action="store_true", + help="Use db_path() from env (AIA_ASSISTANT_DB_PATH / default) as SQLite source", + ) + ap.add_argument( + "--load-system-env", + action="store_true", + help="Merge _local/system.env into the process (for DB_PATH / DATABASE_URL on devices)", + ) + ap.add_argument( + "--pg-url", + default=None, + help="Target PostgreSQL URL; if omitted, use AIA_ASSISTANT_DATABASE_URL (or OPS_* / *_PG_DSN)", + ) + ap.add_argument("--dry-run", action="store_true", help="Count rows only; do not write to PG") + ap.add_argument( + "--allow-non-empty", + action="store_true", + help="Do not abort when target PG tables already contain rows", + ) + ap.add_argument( + "--batch", + type=int, + default=500, + help="Rows per executemany batch (default 500)", + ) + args = ap.parse_args() + if bool(args.load_system_env): + from svc.config.bootstrap_env import load_system_env + + load_system_env(force=True) + + sqlite_path = _resolve_sqlite_path(args.sqlite, bool(args.sqlite_from_db_path)) + if not sqlite_path.is_file(): + raise SystemExit(f"sqlite file not found: {sqlite_path}") + + raw_pg = (args.pg_url or "").strip() or _pg_url_from_environ() + if not raw_pg: + raise SystemExit( + "No PostgreSQL URL: pass --pg-url or set one of " + "AIA_ASSISTANT_DATABASE_URL, OPS_ASSISTANT_DATABASE_URL, " + "AIA_ASSISTANT_PG_DSN, OPS_ASSISTANT_PG_DSN (use --load-system-env to read _local/system.env)." + ) + pg_url = normalize_psycopg_conninfo(raw_pg) + sl = sqlite3.connect(str(sqlite_path)) + sl.row_factory = sqlite3.Row + try: + sl.execute("PRAGMA foreign_keys=ON") + except Exception: + pass + pg = psycopg.connect(pg_url, row_factory=dict_row, autocommit=False) + try: + sqlite_tables = _sqlite_user_tables(sl) + pg_tables = _pg_public_tables(pg) + common = [t for t in sqlite_tables if t in pg_tables] + skipped_sqlite = [t for t in sqlite_tables if t not in pg_tables] + if skipped_sqlite: + print("skip (not in PG public schema):", ", ".join(skipped_sqlite)) + + if not common: + raise SystemExit("No common tables between SQLite and PostgreSQL; nothing to copy.") + + order = _migration_order(sl, common) + if not args.dry_run and not args.allow_non_empty: + _assert_pg_tables_empty(pg, common) + + src_counts = _sqlite_row_counts(sl, order) + total = 0 + copied: list[str] = [] + for t in order: + cols = _common_columns(sl, pg, t) + if not cols: + n0 = src_counts.get(t, 0) + if n0 > 0: + print(f"{t}: SKIP (no common columns; sqlite has {n0} rows — schema drift)") + continue + n = _copy_table(sl=sl, pg=pg, table=t, cols=cols, dry_run=bool(args.dry_run), batch=int(args.batch)) + print(f"{t}: {n} rows ({len(cols)} columns)") + total += n + copied.append(t) + if not args.dry_run: + pg.commit() + _sync_sequences(pg) + pg.commit() + verify = _pg_row_counts(pg, order) + bad = [t for t in copied if verify.get(t, 0) != src_counts.get(t, 0)] + if bad: + print("WARNING: row count mismatch PG vs SQLite for:", ", ".join(bad)) + for t in bad: + print(f" {t}: sqlite={src_counts.get(t, 0)} pg={verify.get(t, 0)}") + else: + print("verify: row counts match SQLite for all copied tables") + print("total rows (copied or dry-run counted):", total) + except BaseException: + pg.rollback() + raise + finally: + sl.close() + pg.close() + + +if __name__ == "__main__": + main() diff --git a/runtime/operations/scripts/seed_mcp_registry.py b/runtime/operations/scripts/seed_mcp_registry.py index 9f57ccd6..12042123 100644 --- a/runtime/operations/scripts/seed_mcp_registry.py +++ b/runtime/operations/scripts/seed_mcp_registry.py @@ -17,6 +17,7 @@ sys.path.insert(0, str(ROOT)) from svc.config.paths import PROJECT_ROOT, db_path # noqa: E402 from svc.persistence.sqlite_store import SqliteStore # noqa: E402 +from svc.persistence.assistant_store import get_assistant_store from runtime.tools.mcp.installer import McpServerManifest, _safe_server_id, install_mcp_server # noqa: E402 @@ -38,7 +39,7 @@ def main(argv: list[str]) -> int: print("no servers[] in seed file", file=sys.stderr) return 2 - store = SqliteStore(db_path()) + store = get_assistant_store() ok_n = 0 for payload in items: if not isinstance(payload, dict): diff --git a/runtime/operations/scripts/set_wecom_auto_bind.py b/runtime/operations/scripts/set_wecom_auto_bind.py index 888e0c0a..5016584d 100644 --- a/runtime/operations/scripts/set_wecom_auto_bind.py +++ b/runtime/operations/scripts/set_wecom_auto_bind.py @@ -4,6 +4,7 @@ import sys from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def _to_bool(v: str) -> str: @@ -20,7 +21,7 @@ def main() -> int: return 2 cmd = args[0].lower() - store = SqliteStore(db_path()) + store = get_assistant_store() if cmd == "show": enabled = str(store.get_setting("wecom_auto_bind_enabled") or "1").strip() diff --git a/runtime/operations/scripts/set_wecom_config.py b/runtime/operations/scripts/set_wecom_config.py index c9c804d8..fb448e55 100644 --- a/runtime/operations/scripts/set_wecom_config.py +++ b/runtime/operations/scripts/set_wecom_config.py @@ -4,6 +4,7 @@ import sys from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def main() -> int: @@ -13,7 +14,7 @@ def main() -> int: print(" python -m scripts.set_wecom_config ") return 2 - store = SqliteStore(db_path()) + store = get_assistant_store() if len(args) < 2: print("error=missing_required_args") return 2 diff --git a/runtime/operations/scripts/smoke_admin_chat_postgres.py b/runtime/operations/scripts/smoke_admin_chat_postgres.py new file mode 100644 index 00000000..39b9dc8d --- /dev/null +++ b/runtime/operations/scripts/smoke_admin_chat_postgres.py @@ -0,0 +1,178 @@ +#!/usr/bin/env python3 +"""HTTP 冒烟:加载 _local/system.env → 用当前配置的 PG → 登录 → 建会话 → POST 一条消息 → 拉历史。 + +不依赖外网 LLM:强制 ``AIA_ASSISTANT_MODE=rule``(本地规则回复)。 + +用法(仓库根):: + + python runtime/operations/scripts/smoke_admin_chat_postgres.py + +退出码 0 表示全流程成功且 ``messages`` 里同时有 user 与 assistant。 +""" + +from __future__ import annotations + +import hashlib +import os +import shutil +import sys +import tempfile +import uuid +from pathlib import Path + + +def _repo_root() -> Path: + return Path(__file__).resolve().parents[3] + + +def main() -> int: + root = _repo_root() + os.chdir(root) + if str(root) not in sys.path: + sys.path.insert(0, str(root)) + + from svc.config.bootstrap_env import load_system_env + + load_system_env(force=True) + + if (os.getenv("AIA_ASSISTANT_DB_BACKEND") or "").strip().lower() not in ( + "postgresql", + "pg", + "postgres", + ): + print("AIA_ASSISTANT_DB_BACKEND 不是 postgresql,本脚本用于测 PG。", file=sys.stderr) + return 2 + + # 不覆盖已有 shell 变量;仅保证规则模式(免外网) + os.environ.setdefault("AIA_ASSISTANT_MODE", "rule") + os.environ["AIA_ASSISTANT_MODE"] = "rule" + + tmp = tempfile.mkdtemp(prefix="oclaw-smoke-ws-") + os.environ["OPS_WORKSPACE_ROOT"] = tmp + + from svc.config import database as db_cfg + from svc.persistence.assistant_store import get_assistant_store, reset_assistant_store_singleton + from svc.persistence.db.engine import clear_assistant_engine_cache + + clear_assistant_engine_cache() + reset_assistant_store_singleton() + + if db_cfg.assistant_db_backend() != "postgresql": + print("assistant_db_backend() 不是 postgresql", file=sys.stderr) + return 2 + + from svc.persistence.sqlite_store import ( + LLM_BUILTIN_RULE_PROFILE_ID, + SqliteStore, + active_llm_profile_setting_key, + ) + + store = get_assistant_store() + assert isinstance(store, SqliteStore) and store._use_pg + + tag = uuid.uuid4().hex[:10] + pw = f"smoke-{tag}" + t = store.create_tenant(f"pg-smoke-{tag}") + tenant_id = str(t["id"]) + # 不用 administrator:共享库里会走全局模型池与已导入的 OpenAI profile,易在无 Key 时得到空 reply。 + u = store.create_user_account( + tenant_id=tenant_id, + username=f"smoke_{tag}", + display_name="Smoke", + role="owner", + password_hash=hashlib.sha256(pw.encode("utf-8")).hexdigest(), + is_active=True, + ) + user_id = str(u.get("id") or "").strip() + if not user_id: + print("create_user_account returned no id", u, file=sys.stderr) + return 1 + store.grant_llm_profile_to_user( + tenant_id=tenant_id, + profile_id=LLM_BUILTIN_RULE_PROFILE_ID, + user_id=user_id, + ) + store.set_setting( + active_llm_profile_setting_key(user_id, f"smoke_{tag}"), + LLM_BUILTIN_RULE_PROFILE_ID, + ) + + from fastapi.testclient import TestClient + + from interfaces.http.fastapi_app import create_app + + client = TestClient(create_app()) + try: + client.post("/admin/api/auth/bootstrap", json={}) + lr = client.post( + "/admin/api/auth/login", + json={ + "tenant_id": tenant_id, + "username": f"smoke_{tag}", + "password": pw, + "purpose": "console", + }, + ) + lj = lr.json() + if not lj.get("ok"): + print("login failed:", lr.status_code, lj, file=sys.stderr) + return 1 + token = str(lj.get("token") or "") + h = {"authorization": f"Bearer {token}"} + + cr = client.post("/admin/api/chat/sessions", json={"title": f"smoke-{tag}"}, headers=h) + if cr.status_code != 200: + print("create session:", cr.status_code, cr.text, file=sys.stderr) + return 1 + cj = cr.json() + if not cj.get("ok"): + print("create session body:", cj, file=sys.stderr) + return 1 + sid = str((cj.get("session") or {}).get("id") or "") + if not sid: + print("no session id", cj, file=sys.stderr) + return 1 + + mr = client.post( + f"/admin/api/chat/sessions/{sid}/messages", + json={"text": "ping-smoke-pg"}, + headers=h, + ) + if mr.status_code != 200: + print("send message:", mr.status_code, mr.text, file=sys.stderr) + return 1 + mj = mr.json() + if not mj.get("ok"): + print("send message body:", mj, file=sys.stderr) + return 1 + reply = str(mj.get("reply") or "") + if not reply.strip(): + print("empty reply (unexpected for rule mode)", mj, file=sys.stderr) + return 1 + + gr = client.get(f"/admin/api/chat/sessions/{sid}/messages", headers=h) + if gr.status_code != 200: + print("list messages:", gr.status_code, gr.text, file=sys.stderr) + return 1 + gj = gr.json() + msgs = gj.get("messages") or [] + roles = [str(m.get("role") or "") for m in msgs if isinstance(m, dict)] + if "user" not in roles or "assistant" not in roles: + print("roles mismatch:", roles, "full ok=", gj.get("ok"), file=sys.stderr) + return 1 + + print("OK smoke_admin_chat_postgres") + print(" tenant_id=", tenant_id) + print(" session_id=", sid) + print(" reply_len=", len(reply)) + print(" message_count=", len(msgs), "roles=", roles) + return 0 + finally: + try: + shutil.rmtree(tmp, ignore_errors=True) + except Exception: + pass + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/runtime/operations/scripts/start_gateway.ps1 b/runtime/operations/scripts/start_gateway.ps1 index 6c6b9b43..0e732d86 100644 --- a/runtime/operations/scripts/start_gateway.ps1 +++ b/runtime/operations/scripts/start_gateway.ps1 @@ -3,7 +3,9 @@ [int]$Port = 8787, [switch]$SkipInstall = $false, [switch]$Background = $false, - [bool]$WithWikiWorker = $true + [bool]$WithWikiWorker = $true, + # Foreground: mirror merged stdout/stderr to gateway.foreground.log (same dir as start_service logs). + [switch]$NoLogMirror = $false ) $ErrorActionPreference = "Stop" @@ -75,6 +77,18 @@ if (-not $SkipInstall) { Write-Step "Skip dependency install" } +# Same directory as runtime.operations.runtime.start_service (db_path parent / logs, or AIA_RUNTIME_LOG_DIR). +$logDir = $null +try { + $logDir = (& $pythonExe -c "from runtime.operations.runtime import assistant_runtime_log_dir; print(str(assistant_runtime_log_dir()))" 2>$null | Select-Object -Last 1).Trim() +} catch { } +if (-not $logDir) { + $logDir = Join-Path $repoRoot "data\logs" +} +Write-Step "Runtime log dir: $logDir" +Write-Host " (stack up / start_service: gateway.err.log + gateway.out.log here)" -ForegroundColor DarkGray +Write-Host " (this script foreground: gateway.foreground.log when log mirror on)" -ForegroundColor DarkGray + Write-Host "" Write-Step "Gateway URL" Write-Host "Admin: http://$BindHost`:$Port/admin" @@ -95,7 +109,15 @@ if ($WithWikiWorker) { if ($Background) { Write-Step "Starting gateway in background" - $p = Start-Process -FilePath $pythonExe -ArgumentList @("-m","runtime.operations","gateway","start","--host",$BindHost,"--port",$Port) -WorkingDirectory $repoRoot -PassThru -WindowStyle Hidden + New-Item -ItemType Directory -Force -Path $logDir | Out-Null + $errLog = Join-Path $logDir "gateway.err.log" + $outLog = Join-Path $logDir "gateway.out.log" + Write-Host "stderr -> $errLog" -ForegroundColor DarkGray + Write-Host "stdout -> $outLog" -ForegroundColor DarkGray + $p = Start-Process -FilePath $pythonExe ` + -ArgumentList @("-m","runtime.operations","gateway","start","--host",$BindHost,"--port",$Port) ` + -WorkingDirectory $repoRoot -PassThru -WindowStyle Hidden ` + -RedirectStandardError $errLog -RedirectStandardOutput $outLog Set-Content -Path $pidFile -Value "$($p.Id)" -Encoding ascii Write-Host "gateway.pid = $pidFile" -ForegroundColor DarkGray Write-Host "PID = $($p.Id)" -ForegroundColor Green @@ -103,7 +125,21 @@ if ($Background) { } Write-Step "Starting gateway (foreground)" -& $pythonExe -m runtime.operations gateway start --host $BindHost --port $Port +New-Item -ItemType Directory -Force -Path $logDir | Out-Null +# Uvicorn writes INFO to stderr. PowerShell 7 wraps native stderr as ErrorRecord +# (red "NativeCommandError") even when non-terminating. Run under cmd.exe with +# ``2>&1`` so PS only sees a single stdout stream of plain text. +if ($BindHost -match '[&|`~]' -or $Port -lt 1 -or $Port -gt 65535) { + Fail "Invalid -BindHost or -Port for gateway launcher." +} +$cmdLine = "cd /d `"$repoRoot`" && `"$pythonExe`" -m runtime.operations gateway start --host $BindHost --port $Port 2>&1" +if (-not $NoLogMirror) { + $fgLog = Join-Path $logDir "gateway.foreground.log" + Write-Host "Mirroring console to: $fgLog (pass -NoLogMirror to disable)" -ForegroundColor DarkGray + cmd.exe /d /s /c $cmdLine | Tee-Object -FilePath $fgLog -Append +} else { + cmd.exe /d /s /c $cmdLine +} diff --git a/runtime/operations/scripts/switch_wecom_bot.py b/runtime/operations/scripts/switch_wecom_bot.py index 0e55d39f..ff116341 100644 --- a/runtime/operations/scripts/switch_wecom_bot.py +++ b/runtime/operations/scripts/switch_wecom_bot.py @@ -4,6 +4,7 @@ import sys from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store _CLEAR_KEYS = [ @@ -29,7 +30,7 @@ def main() -> int: print("error=missing_required_args") return 2 - store = SqliteStore(db_path()) + store = get_assistant_store() store.set_setting("wecom_mode", "bot_api") store.set_setting("wecom_bot_id", bot_id) store.set_secret("wecom_bot_secret", bot_secret) diff --git a/runtime/operations/scripts/unbind_wecom_bot.py b/runtime/operations/scripts/unbind_wecom_bot.py index 01b10625..0263d206 100644 --- a/runtime/operations/scripts/unbind_wecom_bot.py +++ b/runtime/operations/scripts/unbind_wecom_bot.py @@ -2,6 +2,7 @@ from __future__ import annotations from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store _CLEAR_KEYS = [ @@ -24,7 +25,7 @@ _CLEAR_KEYS = [ def main() -> int: - store = SqliteStore(db_path()) + store = get_assistant_store() with store._connect() as conn: # internal cleanup script; safe to use store connection helper cur_ident = conn.execute("DELETE FROM channel_identity WHERE channel = ?", ("wecom",)) ident_deleted = int(cur_ident.rowcount or 0) diff --git a/runtime/operations/scripts/wecom_smoke_test.py b/runtime/operations/scripts/wecom_smoke_test.py index d42a8f17..7a7223cf 100644 --- a/runtime/operations/scripts/wecom_smoke_test.py +++ b/runtime/operations/scripts/wecom_smoke_test.py @@ -7,10 +7,11 @@ import uuid from svc.config.paths import db_path from svc.integrations.wecom_client import WeComClient from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def main() -> int: - store = SqliteStore(db_path()) + store = get_assistant_store() client = WeComClient(store) try: bot_id, _bot_secret = client.get_bot_credentials() diff --git a/runtime/operations/scripts/wecom_status.py b/runtime/operations/scripts/wecom_status.py index e9bb7e96..ea717a29 100644 --- a/runtime/operations/scripts/wecom_status.py +++ b/runtime/operations/scripts/wecom_status.py @@ -6,6 +6,7 @@ from datetime import datetime, timezone from svc.config.paths import db_path from svc.integrations.wecom_client import WeComClient from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def _fmt_ts(ts: str) -> str: @@ -17,7 +18,7 @@ def _fmt_ts(ts: str) -> str: def main() -> int: - store = SqliteStore(db_path()) + store = get_assistant_store() client = WeComClient(store) print("ok=1") print(f"mode={client.mode()}") diff --git a/runtime/operations/scripts/wiki_auto_smoke_test.py b/runtime/operations/scripts/wiki_auto_smoke_test.py index a6cced56..1ab4ffcd 100644 --- a/runtime/operations/scripts/wiki_auto_smoke_test.py +++ b/runtime/operations/scripts/wiki_auto_smoke_test.py @@ -13,6 +13,7 @@ if str(_REPO_ROOT) not in sys.path: from svc.config.paths import PROJECT_ROOT, db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store def _load_cfg() -> dict: @@ -44,7 +45,7 @@ def main() -> int: cfg = _load_cfg() wiki_root = _wiki_root_from_cfg(cfg) - store = SqliteStore(db_path()) + store = get_assistant_store() suffix = uuid.uuid4().hex[:8] payload = { diff --git a/runtime/plan_agent_v2/compat.py b/runtime/plan_agent_v2/compat.py index e5700f5a..58f1e9b6 100644 --- a/runtime/plan_agent_v2/compat.py +++ b/runtime/plan_agent_v2/compat.py @@ -38,6 +38,7 @@ def build_shadow_gateway_result( "relay_ttl_turn_count": 0, "relay_ttl_session_count": 0, "relay_ttl_keep_count": 0, + "turn_uuid": "", } diff --git a/runtime/prompt_prebuild.py b/runtime/prompt_prebuild.py index 6eb2ff71..2f69931d 100644 --- a/runtime/prompt_prebuild.py +++ b/runtime/prompt_prebuild.py @@ -4,8 +4,7 @@ import threading import time from typing import Any -from svc.config.paths import db_path -from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.agent_context import build_role_system_context from runtime.agents.specialists import discover_specialist_ids from runtime.direct_loop import tool_wire_freeze_status, warm_tool_wire_cache @@ -179,7 +178,7 @@ def run_runtime_prewarm( "error": "", } t0 = time.perf_counter() - own_store = store if store is not None else SqliteStore(db_path()) + own_store = store if store is not None else get_assistant_store() try: registry = default_registry(store=own_store) prompt_stats = warm_startup_prompt_prebuild( @@ -243,7 +242,7 @@ def runtime_prewarm_prompts_snapshot( base_url: str = "", memory_enabled: bool = True, ) -> dict[str, Any]: - own_store = store if store is not None else SqliteStore(db_path()) + own_store = store if store is not None else get_assistant_store() registry = default_registry(store=own_store) target = str(role or "").strip().lower() allowed_roles = ["manager", *list(discover_specialist_ids())] diff --git a/runtime/tools/evals/assistant_runner.py b/runtime/tools/evals/assistant_runner.py index d79caaeb..bff55c0a 100644 --- a/runtime/tools/evals/assistant_runner.py +++ b/runtime/tools/evals/assistant_runner.py @@ -8,6 +8,7 @@ from typing import Any from runtime.application.gateway import process_inbound_payload_usecase from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store @dataclass(frozen=True) @@ -59,7 +60,7 @@ def _extract_reply_text(resp: dict[str, Any]) -> str: def run_gateway_eval(dataset_path: str) -> dict[str, Any]: - store = SqliteStore(db_path()) + store = get_assistant_store() # Seed a tenant + bind code for tests tenants = store.list_tenants(limit=1) if tenants: diff --git a/runtime/tools/evals/runner.py b/runtime/tools/evals/runner.py index 559fbf1b..eb82ee0b 100644 --- a/runtime/tools/evals/runner.py +++ b/runtime/tools/evals/runner.py @@ -8,6 +8,7 @@ from typing import Any from runtime.agents.factory import build_gateway_executor from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.orchestration.evaluation import eval_summary from svc.config.paths import db_path from runtime.gateway import OclawGateway @@ -71,7 +72,7 @@ def run_eval( Dataset format: JSONL, each line: {"id": "...", "input": "...", "assert_contains": ["..."], "assert_not_contains": ["..."]} """ - store = SqliteStore(db_path()) + store = get_assistant_store() agent = build_gateway_executor(store) session = store.create_session("offline-eval") gw = OclawGateway(store=store) diff --git a/runtime/tools/experts/productivity/kb_tools.py b/runtime/tools/experts/productivity/kb_tools.py index 915f233e..0497564d 100644 --- a/runtime/tools/experts/productivity/kb_tools.py +++ b/runtime/tools/experts/productivity/kb_tools.py @@ -6,6 +6,7 @@ from typing import Any from svc.config.paths import db_path from svc.embeddings.embedding_client import build_default_embedding_client from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.tools.base import ToolSpec @@ -27,7 +28,7 @@ def kb_add_tool() -> ToolSpec: if title: source = f"{source}:{title[:48]}" cid = _chunk_id(source, text) - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_knowledge_chunk( chunk_id=cid, source=source, @@ -63,7 +64,7 @@ def kb_search_tool() -> ToolSpec: limit = int(args.get("limit") or 3) if not tenant_id or not query: return {"ok": False, "error": "tenant_id and query are required"} - store = SqliteStore(db_path()) + store = get_assistant_store() from runtime.orchestration.memory import retrieve_context rows = retrieve_context(store, query, limit=max(1, min(limit, 6))) diff --git a/runtime/tools/experts/productivity/todo_tools.py b/runtime/tools/experts/productivity/todo_tools.py index 5e4677e0..89f9d90a 100644 --- a/runtime/tools/experts/productivity/todo_tools.py +++ b/runtime/tools/experts/productivity/todo_tools.py @@ -4,6 +4,7 @@ from typing import Any from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.tools.base import ToolSpec @@ -22,7 +23,7 @@ def todo_create_tool() -> ToolSpec: title = _require(str(args.get("title") or ""), "title") due_at = str(args.get("due_at") or "").strip() or None assignee_user_id = str(args.get("assignee_user_id") or "").strip() or None - store = SqliteStore(db_path()) + store = get_assistant_store() row = store.todo_create( tenant_id=tenant_id, owner_user_id=owner_user_id, @@ -61,7 +62,7 @@ def todo_list_tool() -> ToolSpec: assignee_user_id = str(args.get("assignee_user_id") or "").strip() or None status = str(args.get("status") or "open").strip() or None limit = int(args.get("limit") or 50) - store = SqliteStore(db_path()) + store = get_assistant_store() rows = store.todo_list( tenant_id=tenant_id, assignee_user_id=assignee_user_id, @@ -96,7 +97,7 @@ def todo_done_tool() -> ToolSpec: try: tenant_id = _require(str(args.get("tenant_id") or ""), "tenant_id") todo_id = _require(str(args.get("todo_id") or ""), "todo_id") - store = SqliteStore(db_path()) + store = get_assistant_store() ok = store.todo_set_status(tenant_id=tenant_id, todo_id=todo_id, status="done") return {"ok": bool(ok), "todo_id": todo_id} except Exception as e: @@ -122,7 +123,7 @@ def todo_assign_tool() -> ToolSpec: tenant_id = _require(str(args.get("tenant_id") or ""), "tenant_id") todo_id = _require(str(args.get("todo_id") or ""), "todo_id") assignee_user_id = _require(str(args.get("assignee_user_id") or ""), "assignee_user_id") - store = SqliteStore(db_path()) + store = get_assistant_store() ok = store.todo_assign(tenant_id=tenant_id, todo_id=todo_id, assignee_user_id=assignee_user_id) return {"ok": bool(ok), "todo_id": todo_id, "assignee_user_id": assignee_user_id} except Exception as e: diff --git a/runtime/tools/public/index_workspace_tool.py b/runtime/tools/public/index_workspace_tool.py index 175c759f..5f53bce1 100644 --- a/runtime/tools/public/index_workspace_tool.py +++ b/runtime/tools/public/index_workspace_tool.py @@ -3,17 +3,16 @@ from __future__ import annotations from typing import Any from runtime.tools.base import ToolSpec +from svc.persistence.assistant_store import get_assistant_store def index_workspace_tool() -> ToolSpec: def handler(args: dict[str, Any]) -> dict[str, Any]: max_files = int(args.get("max_files") or 120) try: - from svc.persistence.sqlite_store import SqliteStore - from svc.config.paths import db_path from runtime.tools.workspace_indexer import index_workspace - store = SqliteStore(db_path()) + store = get_assistant_store() st = index_workspace(store, max_files=max(1, min(max_files, 800))) return {"ok": True, "files_seen": st.files_seen, "chunks_upserted": st.chunks_upserted, "embeddings_upserted": st.embeddings_upserted} except Exception as e: diff --git a/runtime/tools/public/skill_auto_install_tool.py b/runtime/tools/public/skill_auto_install_tool.py index 72f5957c..25d6bb01 100644 --- a/runtime/tools/public/skill_auto_install_tool.py +++ b/runtime/tools/public/skill_auto_install_tool.py @@ -5,6 +5,7 @@ from typing import Any from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.chat.tool_invocation_context import ( current_tool_lane_sessions, current_tool_workspace_lane_role, @@ -16,7 +17,7 @@ from runtime.tools.base import ToolSpec def _store() -> SqliteStore: - return SqliteStore(db_path()) + return get_assistant_store() def _coerce_public_flag(raw: Any) -> bool: diff --git a/runtime/tools/public/skills_install_tool.py b/runtime/tools/public/skills_install_tool.py index a2cb699f..520433c3 100644 --- a/runtime/tools/public/skills_install_tool.py +++ b/runtime/tools/public/skills_install_tool.py @@ -5,6 +5,7 @@ from typing import Any from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store from runtime.skill_installer import install_skill_from_registry_archive from runtime.skills import default_skills_root from runtime.skills_market import get_market_adapter, normalize_skill_market_provider_setting @@ -12,7 +13,7 @@ from runtime.tools.base import ToolSpec def _store() -> SqliteStore: - return SqliteStore(db_path()) + return get_assistant_store() def _agent_workspace_skills_root() -> Path: diff --git a/runtime/workers/wiki/main.py b/runtime/workers/wiki/main.py index 3c385f50..ea1aa315 100644 --- a/runtime/workers/wiki/main.py +++ b/runtime/workers/wiki/main.py @@ -13,6 +13,7 @@ from typing import Any from svc.config.paths import PROJECT_ROOT, db_path from svc.persistence.sqlite_store import OclawTask, SqliteStore +from svc.persistence.assistant_store import get_assistant_store def _load_oclaw_config() -> dict[str, Any]: @@ -312,7 +313,7 @@ def _process_task(*, task: OclawTask, handlers: dict[str, Any], plugin_cfg: dict def run_worker() -> int: interval_s = max(2, min(int(os.getenv("AIA_WIKI_WORKER_POLL_SECONDS", "4")), 60)) worker_id = str(os.getenv("AIA_WIKI_WORKER_ID") or "wiki-worker-main").strip() or "wiki-worker-main" - store = SqliteStore(db_path()) + store = get_assistant_store() while True: cfg = _load_oclaw_config() plugin_cfg = _wiki_plugin_config(cfg) diff --git a/scripts/clear_postgres_chat_sessions.ps1 b/scripts/clear_postgres_chat_sessions.ps1 new file mode 100644 index 00000000..3c7a15b6 --- /dev/null +++ b/scripts/clear_postgres_chat_sessions.ps1 @@ -0,0 +1,4 @@ +$ErrorActionPreference = "Stop" +$real = Join-Path $PSScriptRoot "..\runtime\operations\scripts\clear_postgres_chat_sessions.ps1" +if (-not (Test-Path $real)) { throw "Forward target not found: $real" } +& $real @args diff --git a/scripts/start_gateway.ps1 b/scripts/start_gateway.ps1 index b1bd5b14..5aa3c612 100644 --- a/scripts/start_gateway.ps1 +++ b/scripts/start_gateway.ps1 @@ -3,7 +3,8 @@ param( [int]$Port = 8787, [switch]$SkipInstall = $false, [switch]$Background = $false, - [bool]$WithWikiWorker = $true + [bool]$WithWikiWorker = $true, + [switch]$NoLogMirror = $false ) $ErrorActionPreference = "Stop" @@ -13,5 +14,5 @@ if (-not (Test-Path $real)) { throw "Forward script target not found: $real" } -& $real -BindHost $BindHost -Port $Port -SkipInstall:$SkipInstall -Background:$Background -WithWikiWorker:$WithWikiWorker +& $real -BindHost $BindHost -Port $Port -SkipInstall:$SkipInstall -Background:$Background -WithWikiWorker:$WithWikiWorker -NoLogMirror:$NoLogMirror diff --git a/svc/config/database.py b/svc/config/database.py new file mode 100644 index 00000000..43f35e29 --- /dev/null +++ b/svc/config/database.py @@ -0,0 +1,57 @@ +"""Assistant database backend selection (SQLite default, PostgreSQL opt-in via env).""" + +from __future__ import annotations + +import os + + +def assistant_db_backend() -> str: + """Return ``sqlite`` (default) or ``postgresql``.""" + raw = ( + os.getenv("AIA_ASSISTANT_DB_BACKEND") + or os.getenv("OPS_ASSISTANT_DB_BACKEND") + or "sqlite" + ).strip().lower() + if raw in ("sqlite", ""): + return "sqlite" + if raw in ("pg", "postgres", "postgresql"): + return "postgresql" + raise ValueError( + f"Invalid assistant DB backend {raw!r}. " + "Use sqlite (default) or postgresql (aliases: pg, postgres)." + ) + + +def assistant_sqlalchemy_url() -> str: + """SQLAlchemy URL for the assistant store (sqlite or postgresql+psycopg).""" + if assistant_db_backend() == "postgresql": + raw = assistant_postgres_dsn() + if raw.startswith("postgresql+") or raw.startswith("postgres+"): + return raw + if raw.startswith("postgresql://") or raw.startswith("postgres://"): + return "postgresql+psycopg://" + raw.split("://", 1)[1] + return raw + from svc.config.paths import db_path + + p = db_path().replace("\\", "/") + return f"sqlite+pysqlite:///{p}" + + +def assistant_postgres_dsn() -> str: + """PostgreSQL connection URI for the assistant store (psycopg/libpq format).""" + url = ( + os.getenv("AIA_ASSISTANT_DATABASE_URL") + or os.getenv("OPS_ASSISTANT_DATABASE_URL") + or os.getenv("AIA_ASSISTANT_PG_DSN") + or os.getenv("OPS_ASSISTANT_PG_DSN") + or "" + ).strip() + if not url: + raise ValueError( + "PostgreSQL backend requires AIA_ASSISTANT_DATABASE_URL (or OPS_ASSISTANT_DATABASE_URL) " + "to a libpq connection string, e.g. postgresql://user:pass@127.0.0.1:5432/oclaw" + ) + return url + + +__all__ = ["assistant_db_backend", "assistant_postgres_dsn", "assistant_sqlalchemy_url"] diff --git a/svc/llm/tool_wire_policy.py b/svc/llm/tool_wire_policy.py index 1065abae..d4f5b691 100644 --- a/svc/llm/tool_wire_policy.py +++ b/svc/llm/tool_wire_policy.py @@ -13,6 +13,7 @@ from datetime import datetime, timedelta, timezone from typing import Any from svc.llm.tool_schema import MIN_OPENAI_FUNCTION_PARAMETERS, complete_openai_tools_wire_parameters +from svc.persistence.assistant_store import get_assistant_store logger = logging.getLogger(__name__) @@ -429,7 +430,7 @@ def prepare_openai_tools_for_llm_api( from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore - store = SqliteStore(db_path()) + store = get_assistant_store() # Role-scoped policies if configured; otherwise fall back to global. policies = load_tool_policies_dict_for_role(store, role=str(role or "").strip().lower() or None) admin = load_merged_admin_config(store) diff --git a/svc/llm/transports/openai_chat_completions.py b/svc/llm/transports/openai_chat_completions.py index 593e5baa..d5cf2f02 100644 --- a/svc/llm/transports/openai_chat_completions.py +++ b/svc/llm/transports/openai_chat_completions.py @@ -5,11 +5,18 @@ import logging import os import re import uuid -from typing import Any, Optional from collections.abc import Callable +from typing import Any, Optional from svc.llm.tool_schema import complete_openai_tools_wire_parameters -from svc.llm.transports.base import ChatModel, LLMResponse, LLMToolCall, normalize_image_b64_payload, coerce_thought_signature_for_storage +from svc.llm.transports.base import ( + ChatModel, + LLMResponse, + LLMToolCall, + coerce_thought_signature_for_storage, + normalize_image_b64_payload, +) +from svc.persistence.assistant_store import get_assistant_store logger = logging.getLogger(__name__) @@ -551,12 +558,10 @@ class OpenAIChatModel(ChatModel): kwargs["extra_body"] = extra_body if use_tools: try: - from svc.config.paths import db_path - from svc.persistence.sqlite_store import SqliteStore from runtime.tools.exposure_plan import build_llm_tools_plan plan = build_llm_tools_plan( - store=SqliteStore(db_path()), + store=get_assistant_store(), role="", base_url=self.base_url, max_json_bytes=_default_max_openai_tools_json_bytes(self.base_url), diff --git a/svc/persistence/assistant_store.py b/svc/persistence/assistant_store.py new file mode 100644 index 00000000..0140441f --- /dev/null +++ b/svc/persistence/assistant_store.py @@ -0,0 +1,57 @@ +"""Single entry point for the assistant persistence layer (SQLite or PostgreSQL).""" + +from __future__ import annotations + +from pathlib import Path + +from svc.config.database import assistant_db_backend, assistant_postgres_dsn +from svc.persistence.assistant_store_protocol import AssistantStoreProtocol + +_singleton: AssistantStoreProtocol | None = None +_singleton_key: str | None = None + + +def reset_assistant_store_singleton() -> None: + """Drop the cached :func:`get_assistant_store` instance (tests / engine URL changes).""" + global _singleton, _singleton_key + _singleton = None + _singleton_key = None + + +def get_assistant_store() -> AssistantStoreProtocol: + """Return the process-wide assistant store implementation. + + - Default: SQLite at :func:`svc.config.paths.db_path`. + - ``AIA_ASSISTANT_DB_BACKEND=postgresql`` + DSN: same :class:`~svc.persistence.sqlite_store.SqliteStore` + API over PostgreSQL (schema via Alembic / ``postgresql_bootstrap.sql``). + + The store is **cached per process** for a stable (backend, connection key) so ``SqliteStore.__init__`` + does not re-run PostgreSQL bootstrap / orphan pruning on every HTTP or WS call (which could race + with in-flight writes and make messages disappear after tool rounds). + + Tests should keep constructing ``SqliteStore(path)`` with an explicit file path; production code + should prefer this factory for ``db_path()``-backed instances. + + Return type is :class:`~svc.persistence.assistant_store_protocol.AssistantStoreProtocol`; the + concrete class is :class:`~svc.persistence.sqlite_store.SqliteStore` for both backends. + """ + global _singleton, _singleton_key + from svc.config.paths import db_path + from svc.persistence.sqlite_store import SqliteStore + + if assistant_db_backend() == "postgresql": + key = f"postgresql::{assistant_postgres_dsn()}" + else: + key = f"sqlite::{Path(db_path()).resolve()}" + + if _singleton is not None and _singleton_key == key: + return _singleton + if assistant_db_backend() == "postgresql": + _singleton = SqliteStore(None, postgres_url=assistant_postgres_dsn()) + else: + _singleton = SqliteStore(db_path()) + _singleton_key = key + return _singleton + + +__all__ = ["get_assistant_store", "reset_assistant_store_singleton"] diff --git a/svc/persistence/assistant_store_protocol.py b/svc/persistence/assistant_store_protocol.py new file mode 100644 index 00000000..854ddc5a --- /dev/null +++ b/svc/persistence/assistant_store_protocol.py @@ -0,0 +1,192 @@ +"""Typing protocol for assistant persistence (implemented by SqliteStore). + +When adding or renaming public methods on :class:`~svc.persistence.sqlite_store.SqliteStore`, +update this protocol (e.g. re-run a small ``inspect.signature`` generator over ``SqliteStore``) +so static checkers stay aligned with :func:`~svc.persistence.assistant_store.get_assistant_store`. +""" + +from __future__ import annotations + +from typing import Any, Optional, Protocol + +from svc.persistence.sqlite_store import ( + ChatMessage, + ChatSession, + OclawRun, + OclawTask, + SessionMessagesMeta, + SessionsListMeta, +) + + +class AssistantStoreProtocol(Protocol): + """Structural contract for :func:`~svc.persistence.assistant_store.get_assistant_store`.""" + + _use_pg: bool + db_path: str + def add_admin_audit_log(self, *, actor_tenant_id: 'str', actor_user_id: 'str', action: 'str', target_type: 'str', target_id: 'str', status: 'str', detail: 'dict[str, Any] | None' = None) -> 'None': ... + def add_agent_audit_log(self, *, session_id: 'str', specialist: 'str', task_kind: 'str', action: 'str', payload: 'dict[str, Any]', status: 'str', reason: 'str', duration_ms: 'int' = 0) -> 'None': ... + def add_agent_eval_log(self, *, session_id: 'str', specialist: 'str', task_kind: 'str', success: 'bool', latency_ms: 'int', cost_hint: 'float' = 0.0, notes: 'str' = '') -> 'None': ... + def add_mcp_installation_log(self, *, server_id: 'str', status: 'str', error_code: 'str' = '', detail: 'dict[str, Any] | None' = None, install_command: 'str' = '') -> 'None': ... + def add_memory_hit_log(self, *, tenant_id: 'str', user_id: 'str', session_id: 'str | None', memory_id: 'str | None', query_text: 'str', score: 'float', source: 'str', timestamp: 'str | None' = None) -> 'None': ... + def add_message(self, session_id: 'str', role: 'str', content: 'str', tool_calls: 'Any | None' = None, attachments: 'Any | None' = None, turn_uuid: 'str | None' = None, event_type: 'str | None' = None, event_payload: 'Any | None' = None, timestamp: 'str | None' = None) -> 'ChatMessage': ... + def add_tool_log(self, session_id: 'str', tool_name: 'str', args: 'dict[str, Any]', result: 'Any', specialist: 'str | None' = None, timestamp: 'str | None' = None, duration_ms: 'int | None' = None) -> 'None': ... + def add_trace_event(self, *, session_id: 'str', trace_id: 'str', span_id: 'str', parent_span_id: 'str | None', event_type: 'str', payload: 'dict[str, Any]') -> 'None': ... + def add_trace_events_batch(self, events: 'list[dict[str, Any]]') -> 'None': ... + def attachment_acl_allows_tenant(self, *, tenant_id: 'str', attachment_id: 'str') -> 'bool': ... + def attachment_acl_allows_user(self, *, tenant_id: 'str', user_id: 'str', attachment_id: 'str') -> 'bool': ... + def attachment_referenced_by_user(self, *, tenant_id: 'str', user_id: 'str', attachment_id: 'str', scan_limit: 'int' = 2000) -> 'bool': ... + def attachment_referenced_in_tenant(self, *, tenant_id: 'str', attachment_id: 'str', scan_limit: 'int' = 4000) -> 'bool': ... + def backfill_attachment_acl_from_messages(self, *, tenant_id: 'str', limit_messages: 'int' = 50000) -> 'dict[str, Any]': ... + def backfill_orphan_chat_sessions_for_user(self, *, tenant_id: 'str', user_id: 'str') -> 'int': ... + def backfill_ui_session_owner_from_channel_v2(self) -> 'int': ... + def backfill_user_channel_account_names(self, *, channel: 'str' = 'wecom') -> 'int': ... + def clear_llm_profile_secret(self, profile_id: 'str') -> 'None': ... + def clear_low_confidence_memory(self, *, max_confidence: 'float') -> 'int': ... + def consume_bind_code(self, *, code: 'str', channel: 'str', external_user_id: 'str', display_name: 'str | None' = None) -> 'dict[str, Any] | None': ... + def count_admin_audit_logs(self, *, tenant_id: 'str | None' = None, action: 'str | None' = None, actor_user_id: 'str | None' = None, status: 'str | None' = None) -> 'int': ... + def count_messages(self, session_id: 'str') -> 'int': ... + def count_sessions(self) -> 'int': ... + def create_auth_session(self, *, session_token_hash: 'str', tenant_id: 'str', user_id: 'str', role: 'str', expires_at: 'str') -> 'None': ... + def create_bind_code(self, *, tenant_id: 'str', role: 'str', code: 'str') -> 'dict[str, Any]': ... + def create_llm_profile(self, name: 'str', mode: 'str' = 'openai', model: 'str | None' = None, base_url: 'str | None' = None, *, owner_user_id: 'str | None' = None) -> 'str': ... + def create_session(self, title: 'str') -> 'ChatSession': ... + def create_session_for_user(self, *, title: 'str', tenant_id: 'str', user_id: 'str') -> 'ChatSession': ... + def create_tenant(self, name: 'str') -> 'dict[str, Any]': ... + def create_user(self, *, tenant_id: 'str', display_name: 'str', role: 'str') -> 'dict[str, Any]': ... + def create_user_account(self, *, tenant_id: 'str', username: 'str', display_name: 'str', role: 'str', password_hash: 'str', is_active: 'bool' = True) -> 'dict[str, Any]': ... + def delete_llm_profile(self, profile_id: 'str') -> 'None': ... + def delete_mcp_server(self, *, server_id: 'str') -> 'dict[str, int]': ... + def delete_memory_item(self, *, memory_id: 'str') -> 'int': ... + def delete_message(self, *, session_id: 'str', message_id: 'int') -> 'bool': ... + def delete_session(self, session_id: 'str') -> 'None': ... + def delete_session_for_user(self, *, session_id: 'str', tenant_id: 'str', user_id: 'str') -> 'bool': ... + def delete_session_in_tenant(self, *, session_id: 'str', tenant_id: 'str') -> 'bool': ... + def delete_setting(self, key: 'str') -> 'None': ... + def delete_tenant(self, *, tenant_id: 'str') -> 'int': ... + def delete_user_account(self, *, tenant_id: 'str', user_id: 'str') -> 'int': ... + def delete_user_channel_account(self, *, tenant_id: 'str', user_id: 'str', channel: 'str', account_id: 'str') -> 'int': ... + def delete_user_permission(self, *, tenant_id: 'str', user_id: 'str', permission: 'str') -> 'int': ... + def ensure_default_session(self) -> 'ChatSession': ... + def ensure_memory_tables(self) -> 'None': ... + def ensure_personal_llm_clone_from_global(self, user_id: 'str', username: 'str | None') -> 'None': ... + def ensure_ui_session_owner(self, *, session_id: 'str', tenant_id: 'str', user_id: 'str') -> 'None': ... + def find_user_by_channel_account(self, *, channel: 'str', account_id: 'str') -> 'dict[str, Any] | None': ... + def fork_session(self, source_session_id: 'str', up_to_message_id: 'int', title: 'str') -> 'ChatSession': ... + def get_auth_session(self, *, session_token_hash: 'str') -> 'dict[str, Any] | None': ... + def get_knowledge_chunks(self, *, chunk_ids: 'list[str]') -> 'list[dict[str, Any]]': ... + def get_last_message_id(self, session_id: 'str') -> 'int | None': ... + def get_llm_profile(self, profile_id: 'str') -> 'Optional[dict[str, Any]]': ... + def get_llm_profile_secret(self, profile_id: 'str') -> 'Optional[str]': ... + def get_messages(self, session_id: 'str', limit: 'int' = 200) -> 'list[ChatMessage]': ... + def get_messages_after_id(self, *, session_id: 'str', after_id: 'int', limit: 'int' = 200) -> 'list[ChatMessage]': ... + def get_or_create_channel_session(self, *, tenant_id: 'str', channel: 'str', external_chat_id: 'str', external_user_id: 'str', session_title: 'str') -> 'str': ... + def get_or_create_channel_session_v2(self, *, tenant_id: 'str', channel: 'str', account_id: 'str', external_chat_id: 'str', external_user_id: 'str', session_title: 'str') -> 'str': ... + def get_secret(self, key: 'str') -> 'Optional[str]': ... + def get_session(self, session_id: 'str') -> 'Optional[ChatSession]': ... + def get_session_for_user(self, *, session_id: 'str', tenant_id: 'str', user_id: 'str') -> 'Optional[ChatSession]': ... + def get_session_in_tenant(self, *, session_id: 'str', tenant_id: 'str') -> 'Optional[ChatSession]': ... + def get_session_messages_meta(self, session_id: 'str') -> 'SessionMessagesMeta': ... + def get_sessions_list_meta(self) -> 'SessionsListMeta': ... + def get_sessions_list_meta_for_tenant(self, *, tenant_id: 'str') -> 'SessionsListMeta': ... + def get_sessions_list_meta_for_user(self, *, tenant_id: 'str', user_id: 'str') -> 'SessionsListMeta': ... + def get_setting(self, key: 'str') -> 'Optional[str]': ... + def get_tool_logs(self, session_id: 'str', limit: 'int' = 200) -> 'list[dict[str, Any]]': ... + def get_turn_time_window(self, *, session_id: 'str', trace_id: 'str') -> 'tuple[str | None, str | None]': ... + def get_ui_session_owner(self, *, session_id: 'str') -> 'dict[str, Any] | None': ... + def get_user_by_id(self, *, tenant_id: 'str', user_id: 'str') -> 'dict[str, Any] | None': ... + def get_user_by_username(self, *, tenant_id: 'str', username: 'str') -> 'dict[str, Any] | None': ... + def get_user_by_username_global(self, *, username: 'str') -> 'dict[str, Any] | None': ... + def get_user_workspace_path_allowlist(self, *, tenant_id: 'str', user_id: 'str') -> 'dict[str, Any] | None': ... + def grant_llm_profile_to_tenant(self, *, tenant_id: 'str', profile_id: 'str', created_by_user_id: 'str | None' = None) -> 'str': ... + def grant_llm_profile_to_user(self, *, tenant_id: 'str', profile_id: 'str', user_id: 'str', created_by_user_id: 'str | None' = None) -> 'str': ... + def legacy_secret_stats(self) -> 'dict[str, Any]': ... + def link_attachment_acl(self, *, tenant_id: 'str', user_id: 'str', session_id: 'str', attachment_id: 'str', source: 'str') -> 'None': ... + def list_admin_audit_logs(self, *, tenant_id: 'str | None' = None, action: 'str | None' = None, actor_user_id: 'str | None' = None, status: 'str | None' = None, limit: 'int' = 200, offset: 'int' = 0) -> 'list[dict[str, Any]]': ... + def list_admin_sessions(self, *, tenant_id: 'str', user_id: 'str | None' = None, q: 'str | None' = None, active_only: 'bool' = False, active_window_minutes: 'int' = 30, limit: 'int' = 100, offset: 'int' = 0) -> 'tuple[int, list[dict[str, Any]]]': ... + def list_admin_user_stats(self, *, tenant_id: 'str', q: 'str | None' = None, active_window_minutes: 'int' = 30, limit: 'int' = 100, offset: 'int' = 0) -> 'tuple[int, list[dict[str, Any]], dict[str, Any]]': ... + def list_agent_audit_logs(self, *, limit: 'int' = 200, session_id: 'str | None' = None) -> 'list[dict[str, Any]]': ... + def list_agent_eval_logs(self, *, limit: 'int' = 200) -> 'list[dict[str, Any]]': ... + def list_bind_codes(self, *, tenant_id: 'str | None' = None, limit: 'int' = 200) -> 'list[dict[str, Any]]': ... + def list_channel_identities(self, *, tenant_id: 'str | None' = None, channel: 'str | None' = None, limit: 'int' = 300) -> 'list[dict[str, Any]]': ... + def list_channel_identities_v2(self, *, tenant_id: 'str | None' = None, channel: 'str | None' = None, account_id: 'str | None' = None, user_id: 'str | None' = None, limit: 'int' = 300) -> 'list[dict[str, Any]]': ... + def list_knowledge_embeddings(self, *, model: 'str', limit: 'int' = 5000) -> 'list[dict[str, Any]]': ... + def list_llm_profile_grants_for_profile(self, tenant_id: 'str', profile_id: 'str') -> 'list[dict[str, Any]]': ... + def list_llm_profiles(self, *, visible_only: 'bool' = False, viewer_user_id: 'str | None' = None, viewer_username: 'str | None' = None, viewer_tenant_id: 'str | None' = None) -> 'list[dict[str, Any]]': ... + def list_mcp_install_failure_summary(self, *, limit: 'int' = 20) -> 'list[dict[str, Any]]': ... + def list_mcp_installation_logs(self, *, server_id: 'str | None' = None, limit: 'int' = 100) -> 'list[dict[str, Any]]': ... + def list_mcp_server_health(self) -> 'list[dict[str, Any]]': ... + def list_mcp_server_tools(self, *, server_id: 'str') -> 'list[dict[str, Any]]': ... + def list_mcp_servers(self, *, enabled_only: 'bool' = False) -> 'list[dict[str, Any]]': ... + def list_mcp_tool_aggregate_usage(self) -> 'dict[str, dict[str, Any]]': ... + def list_mcp_tool_call_logs(self, *, server_id: 'str | None' = None, limit: 'int' = 200) -> 'list[dict[str, Any]]': ... + def list_mcp_tool_usage_summary(self, *, limit: 'int' = 200) -> 'list[dict[str, Any]]': ... + def list_memory_hit_logs(self, *, tenant_id: 'str | None' = None, user_id: 'str | None' = None, limit: 'int' = 100) -> 'list[dict[str, Any]]': ... + def list_memory_items(self, *, tenant_id: 'str | None' = None, user_id: 'str | None' = None, limit: 'int' = 100, offset: 'int' = 0) -> 'list[dict[str, Any]]': ... + def list_messages_in_time_window(self, *, session_id: 'str', start_ts: 'str | None', end_ts: 'str | None', limit: 'int' = 500) -> 'list[dict[str, Any]]': ... + def list_session_tool_health(self, *, session_id: 'str | None' = None, limit: 'int' = 80) -> 'list[dict[str, Any]]': ... + def list_sessions(self, limit: 'int | None' = None, offset: 'int' = 0) -> 'list[ChatSession]': ... + def list_sessions_for_tenant(self, *, tenant_id: 'str', limit: 'int | None' = None, offset: 'int' = 0) -> 'list[ChatSession]': ... + def list_sessions_for_user(self, *, tenant_id: 'str', user_id: 'str', limit: 'int | None' = None, offset: 'int' = 0) -> 'list[ChatSession]': ... + def list_tenants(self, *, limit: 'int' = 200) -> 'list[dict[str, Any]]': ... + def list_tool_plugins(self) -> 'list[dict[str, Any]]': ... + def list_trace_events(self, *, session_id: 'str', limit: 'int' = 300) -> 'list[dict[str, Any]]': ... + def list_trace_events_for_trace(self, *, session_id: 'str', trace_id: 'str', limit: 'int' = 500) -> 'list[dict[str, Any]]': ... + def list_user_channel_accounts(self, *, tenant_id: 'str', user_id: 'str', channel: 'str' = 'wecom', include_inactive: 'bool' = True) -> 'list[dict[str, Any]]': ... + def list_user_permissions(self, *, tenant_id: 'str', user_id: 'str', role: 'str | None' = None) -> 'list[str]': ... + def list_user_workspace_extra_roots_union(self) -> 'list[str]': ... + def list_users(self, *, tenant_id: 'str', limit: 'int' = 500, offset: 'int' = 0, q: 'str | None' = None, include_inactive: 'bool' = True) -> 'list[dict[str, Any]]': ... + def migrate_secrets_to_fernet(self) -> 'dict[str, int]': ... + def move_tool_logs_to_session(self, *, from_session_id: 'str', to_session_id: 'str') -> 'int': ... + def oclaw_attempt_append(self, *, run_id: 'str', tenant_id: 'str', session_id: 'str', attempt_no: 'int', status: 'str', reason: 'str' = '', payload: 'dict[str, Any] | None' = None) -> 'int': ... + def oclaw_attempt_list(self, *, run_id: 'str', limit: 'int' = 30) -> 'list[dict[str, Any]]': ... + def oclaw_run_get(self, *, run_id: 'str', tenant_id: 'str | None' = None) -> 'OclawRun | None': ... + def oclaw_run_list(self, *, tenant_id: 'str', session_id: 'str | None' = None, status: 'str | None' = None, limit: 'int' = 50) -> 'list[OclawRun]': ... + def oclaw_run_upsert(self, *, run_id: 'str', tenant_id: 'str', session_id: 'str', status: 'str', payload: 'dict[str, Any] | None' = None) -> 'bool': ... + def oclaw_task_claim(self, *, worker_id: 'str', lease_seconds: 'int' = 90, task_type: 'str | None' = None) -> 'OclawTask | None': ... + def oclaw_task_create(self, *, tenant_id: 'str', session_id: 'str', task_type: 'str' = 'async_turn', payload: 'dict[str, Any] | None' = None) -> 'OclawTask': ... + def oclaw_task_fail(self, *, task_id: 'str', error: 'str', result: 'dict[str, Any] | None' = None) -> 'bool': ... + def oclaw_task_finish(self, *, task_id: 'str', result: 'dict[str, Any] | None' = None) -> 'bool': ... + def oclaw_task_get(self, *, task_id: 'str', tenant_id: 'str | None' = None) -> 'OclawTask | None': ... + def oclaw_task_list(self, *, status: 'str | None' = None, limit: 'int' = 50, tenant_id: 'str | None' = None, session_id: 'str | None' = None) -> 'list[OclawTask]': ... + def rename_session(self, session_id: 'str', title: 'str') -> 'None': ... + def replace_mcp_server_tools(self, *, server_id: 'str', tools: 'list[dict[str, Any]]') -> 'None': ... + def resolve_user_by_channel_identity(self, *, channel: 'str', external_user_id: 'str') -> 'dict[str, Any] | None': ... + def resolve_user_by_channel_identity_v2(self, *, channel: 'str', account_id: 'str', external_user_id: 'str') -> 'dict[str, Any] | None': ... + def revoke_all_auth_sessions(self) -> 'int': ... + def revoke_auth_session(self, *, session_token_hash: 'str') -> 'int': ... + def revoke_llm_profile_grant(self, *, tenant_id: 'str', profile_id: 'str', user_id: 'str') -> 'int': ... + def revoke_llm_profile_tenant_grant(self, *, tenant_id: 'str', profile_id: 'str') -> 'int': ... + def search_knowledge(self, *, query: 'str', limit: 'int' = 3) -> 'list[dict[str, Any]]': ... + def search_memory_vectors(self, *, query_vector: 'list[float]', model: 'str', tenant_id: 'str', user_id: 'str', limit: 'int' = 5) -> 'list[dict[str, Any]]': ... + def set_llm_profile_secret(self, profile_id: 'str', plain_text: 'str') -> 'None': ... + def set_mcp_server_enabled(self, *, server_id: 'str', enabled: 'bool') -> 'int': ... + def set_mcp_server_health(self, *, server_id: 'str', status: 'str', detail: 'dict[str, Any] | None' = None) -> 'None': ... + def set_secret(self, key: 'str', plain_text: 'str') -> 'None': ... + def set_setting(self, key: 'str', value: 'str') -> 'None': ... + def sync_attachment_acl_from_chat_message_attachments(self, *, session_id: 'str', role: 'str', event_type: 'str | None', attachments: 'Any') -> 'None': ... + def tenant_has_llm_profile_grant(self, tenant_id: 'str', profile_id: 'str') -> 'bool': ... + def todo_assign(self, *, tenant_id: 'str', todo_id: 'str', assignee_user_id: 'str') -> 'bool': ... + def todo_create(self, *, tenant_id: 'str', owner_user_id: 'str', title: 'str', due_at: 'str | None' = None, assignee_user_id: 'str | None' = None) -> 'dict[str, Any]': ... + def todo_list(self, *, tenant_id: 'str', assignee_user_id: 'str | None' = None, status: 'str | None' = 'open', limit: 'int' = 50) -> 'list[dict[str, Any]]': ... + def todo_set_status(self, *, tenant_id: 'str', todo_id: 'str', status: 'str') -> 'bool': ... + def touch_auth_session(self, *, session_token_hash: 'str') -> 'None': ... + def trim_messages(self, session_id: 'str', keep_last: 'int') -> 'None': ... + def update_llm_profile(self, profile_id: 'str', name: 'str', mode: 'str', model: 'str | None', base_url: 'str | None', *, thinking_mode_enabled: 'bool | None' = None, reasoning_effort: 'str | None' = None) -> 'None': ... + def update_message_content(self, *, session_id: 'str', message_id: 'int', content: 'str', event_payload: 'Any | None' = None) -> 'bool': ... + def update_user_account(self, *, tenant_id: 'str', user_id: 'str', display_name: 'str | None' = None, role: 'str | None' = None, is_active: 'bool | None' = None, password_hash: 'str | None' = None, avatar_attachment_id: 'str | None' = None) -> 'bool': ... + def upsert_channel_identity(self, *, tenant_id: 'str', channel: 'str', external_user_id: 'str', user_id: 'str') -> 'None': ... + def upsert_channel_identity_v2(self, *, tenant_id: 'str', channel: 'str', account_id: 'str', external_user_id: 'str', user_id: 'str') -> 'None': ... + def upsert_knowledge_chunk(self, *, chunk_id: 'str', source: 'str', content: 'str', metadata: 'dict[str, Any] | None' = None) -> 'None': ... + def upsert_knowledge_embedding(self, *, chunk_id: 'str', model: 'str', vector: 'list[float]') -> 'None': ... + def upsert_mcp_server(self, *, server_id: 'str', source_type: 'str', source_ref: 'str', version: 'str' = '', entry_command: 'str' = '', entry_args: 'list[str] | None' = None, env_schema: 'dict[str, Any] | None' = None, required_permissions: 'list[str] | None' = None, risk_level: 'str' = 'high', timeout_s: 'float' = 30.0, enabled: 'bool' = False) -> 'None': ... + def upsert_memory_item(self, *, memory_id: 'str', tenant_id: 'str', user_id: 'str', session_id: 'str', memory_type: 'str', content: 'str', confidence: 'float', source: 'str', metadata: 'dict[str, Any] | None' = None, created_at: 'str | None' = None, updated_at: 'str | None' = None, expires_at: 'str | None' = None) -> 'None': ... + def upsert_memory_vector(self, *, memory_id: 'str', model: 'str', vector: 'list[float]', updated_at: 'str | None' = None) -> 'None': ... + def upsert_tool_plugin(self, *, plugin_name: 'str', plugin_version: 'str', entry_point: 'str', enabled: 'bool' = True) -> 'None': ... + def upsert_user_channel_account(self, *, tenant_id: 'str', user_id: 'str', channel: 'str', account_id: 'str', name: 'str | None' = None, config: 'dict[str, Any] | None' = None, is_active: 'bool' = True) -> 'None': ... + def upsert_user_permission(self, *, tenant_id: 'str', user_id: 'str', permission: 'str') -> 'None': ... + def upsert_user_workspace_path_allowlist(self, *, tenant_id: 'str', user_id: 'str', extra_roots: 'str', allow_any_path: 'bool', allow_high_risk_public_tools: 'bool' = False) -> 'None': ... + def user_has_llm_profile_grant(self, tenant_id: 'str', user_id: 'str', profile_id: 'str') -> 'bool': ... + + +__all__ = ["AssistantStoreProtocol"] diff --git a/svc/persistence/db/__init__.py b/svc/persistence/db/__init__.py new file mode 100644 index 00000000..eab14059 --- /dev/null +++ b/svc/persistence/db/__init__.py @@ -0,0 +1,43 @@ +"""Persistence DB helpers.""" + +from svc.persistence.db.engine import ( + clear_assistant_engine_cache, + engine_for_sqlite_file, + get_assistant_engine, +) +from svc.persistence.db.tables import ( + app_setting, + app_user, + auth_session, + bind_code, + channel_identity, + channel_identity_v2, + channel_session_v2, + chat_message, + chat_session, + metadata, + tenant, + tool_log, + trace_event, + ui_session_owner, +) + +__all__ = [ + "app_user", + "app_setting", + "auth_session", + "bind_code", + "channel_identity", + "channel_identity_v2", + "channel_session_v2", + "chat_message", + "chat_session", + "clear_assistant_engine_cache", + "engine_for_sqlite_file", + "get_assistant_engine", + "metadata", + "tenant", + "tool_log", + "trace_event", + "ui_session_owner", +] diff --git a/svc/persistence/db/engine.py b/svc/persistence/db/engine.py new file mode 100644 index 00000000..fd9e7dff --- /dev/null +++ b/svc/persistence/db/engine.py @@ -0,0 +1,72 @@ +"""SQLAlchemy Engine factory for assistant DB (SQLite or PostgreSQL).""" + +from __future__ import annotations + +from functools import lru_cache +from pathlib import Path +from typing import Any + +from sqlalchemy import create_engine, event +from sqlalchemy.engine import Engine +from sqlalchemy.pool import NullPool + +from svc.config.database import assistant_sqlalchemy_url + + +def _sqlite_sa_url_from_os_path(path: str) -> str: + """Build the same ``sqlite+pysqlite:///...`` URL shape as :func:`svc.config.database.assistant_sqlalchemy_url`.""" + p = Path(path).resolve().as_posix() + return f"sqlite+pysqlite:///{p}" + + +def _register_sqlite_pragmas(eng: Engine) -> None: + @event.listens_for(eng, "connect") + def _sqlite_pragmas(dbapi_conn: Any, _record: Any) -> None: + cur = dbapi_conn.cursor() + cur.execute("PRAGMA foreign_keys = ON;") + cur.execute("PRAGMA journal_mode = WAL;") + cur.execute("PRAGMA synchronous = NORMAL;") + cur.execute("PRAGMA busy_timeout = 30000;") + cur.close() + + +@lru_cache(maxsize=64) +def _engine_for_url(url: str) -> Engine: + """One engine per URL (process-wide).""" + pool_kw: dict[str, Any] = {} + if url.startswith("sqlite"): + pool_kw["poolclass"] = NullPool + else: + pool_kw["pool_pre_ping"] = True + eng = create_engine(url, future=True, **pool_kw) + if url.startswith("sqlite"): + _register_sqlite_pragmas(eng) + return eng + + +def get_assistant_engine() -> Engine: + """Engine for the current env-selected assistant DB (Alembic, ``get_assistant_store`` SQLite path).""" + return _engine_for_url(assistant_sqlalchemy_url()) + + +def engine_for_sqlite_file(path: str) -> Engine: + """Engine for a specific SQLite file (e.g. ``SqliteStore('/tmp/x.sqlite')`` without env ``DB_PATH``).""" + return _engine_for_url(_sqlite_sa_url_from_os_path(path)) + + +def clear_assistant_engine_cache() -> None: + """Drop cached engines (e.g. tests that delete temp DB files must call this before removing the directory).""" + _engine_for_url.cache_clear() + try: + from svc.persistence.assistant_store import reset_assistant_store_singleton + + reset_assistant_store_singleton() + except Exception: + pass + + +__all__ = [ + "clear_assistant_engine_cache", + "engine_for_sqlite_file", + "get_assistant_engine", +] diff --git a/svc/persistence/db/tables.py b/svc/persistence/db/tables.py new file mode 100644 index 00000000..f9abb983 --- /dev/null +++ b/svc/persistence/db/tables.py @@ -0,0 +1,171 @@ +"""SQLAlchemy Core table objects for incremental persistence migration.""" + +from __future__ import annotations + +from sqlalchemy import BigInteger, Column, ForeignKey, Integer, MetaData, Table, Text + +metadata = MetaData() + +tenant = Table( + "tenant", + metadata, + Column("id", Text, primary_key=True), + Column("name", Text, nullable=False), + Column("created_at", Text, nullable=False), +) + +bind_code = Table( + "bind_code", + metadata, + Column("code", Text, primary_key=True), + Column("tenant_id", Text, nullable=False), + Column("role", Text, nullable=False), + Column("created_at", Text, nullable=False), + Column("used_at", Text, nullable=True), + Column("used_by_external_user_id", Text, nullable=True), +) + +app_user = Table( + "app_user", + metadata, + Column("id", Text, primary_key=True), + Column("tenant_id", Text, nullable=False), + Column("username", Text, nullable=True), + Column("display_name", Text, nullable=False), + Column("role", Text, nullable=False), + Column("password_hash", Text, nullable=True), + Column("is_active", Integer, nullable=False, server_default="1"), + Column("created_at", Text, nullable=False), + Column("avatar_attachment_id", Text, nullable=True), +) + +app_setting = Table( + "app_setting", + metadata, + Column("key", Text, primary_key=True), + Column("value", Text, nullable=False), + Column("is_secret", Integer, nullable=False, server_default="0"), + Column("updated_at", Text, nullable=False), +) + +auth_session = Table( + "auth_session", + metadata, + Column("session_token_hash", Text, primary_key=True), + Column("tenant_id", Text, nullable=False), + Column("user_id", Text, nullable=False), + Column("role", Text, nullable=False), + Column("created_at", Text, nullable=False), + Column("expires_at", Text, nullable=False), + Column("last_seen_at", Text, nullable=False), + Column("revoked_at", Text, nullable=True), +) + +chat_session = Table( + "chat_session", + metadata, + Column("id", Text, primary_key=True), + Column("title", Text, nullable=False), + Column("created_at", Text, nullable=False), + Column("last_message_at", Text, nullable=True), +) + +channel_identity_v2 = Table( + "channel_identity_v2", + metadata, + Column("tenant_id", Text, primary_key=True), + Column("channel", Text, primary_key=True), + Column("account_id", Text, primary_key=True), + Column("external_user_id", Text, primary_key=True), + Column("user_id", Text, nullable=False), + Column("created_at", Text, nullable=False), +) + +channel_identity = Table( + "channel_identity", + metadata, + Column("tenant_id", Text, primary_key=True), + Column("channel", Text, primary_key=True), + Column("external_user_id", Text, primary_key=True), + Column("user_id", Text, nullable=False), + Column("created_at", Text, nullable=False), +) + +channel_session_v2 = Table( + "channel_session_v2", + metadata, + Column("tenant_id", Text, primary_key=True), + Column("channel", Text, primary_key=True), + Column("account_id", Text, primary_key=True), + Column("external_chat_id", Text, primary_key=True), + Column("external_user_id", Text, primary_key=True), + Column("session_id", Text, nullable=False), + Column("created_at", Text, nullable=False), +) + +ui_session_owner = Table( + "ui_session_owner", + metadata, + Column("session_id", Text, primary_key=True), + Column("tenant_id", Text, nullable=False), + Column("user_id", Text, nullable=False), + Column("created_at", Text, nullable=False), +) + +chat_message = Table( + "chat_message", + metadata, + Column("id", BigInteger, primary_key=True, autoincrement=True), + Column("session_id", Text, ForeignKey("chat_session.id", ondelete="CASCADE"), nullable=False), + Column("role", Text, nullable=False), + Column("content", Text, nullable=False), + Column("tool_calls", Text, nullable=True), + Column("attachments", Text, nullable=True), + Column("turn_uuid", Text, nullable=True), + Column("event_type", Text, nullable=True), + Column("event_payload", Text, nullable=True), + Column("timestamp", Text, nullable=False), +) + +tool_log = Table( + "tool_log", + metadata, + Column("id", BigInteger, primary_key=True, autoincrement=True), + Column("session_id", Text, ForeignKey("chat_session.id", ondelete="CASCADE"), nullable=False), + Column("tool_name", Text, nullable=False), + Column("specialist", Text, nullable=False, server_default=""), + Column("args", Text, nullable=False), + Column("result", Text, nullable=False), + Column("timestamp", Text, nullable=False), + Column("duration_ms", Integer, nullable=True), +) + +trace_event = Table( + "trace_event", + metadata, + Column("id", BigInteger, primary_key=True, autoincrement=True), + Column("session_id", Text, nullable=False), + Column("trace_id", Text, nullable=False), + Column("span_id", Text, nullable=False), + Column("parent_span_id", Text, nullable=True), + Column("event_type", Text, nullable=False), + Column("payload", Text, nullable=False, server_default="{}"), + Column("timestamp", Text, nullable=False), +) + +__all__ = [ + "app_user", + "app_setting", + "auth_session", + "bind_code", + "channel_identity", + "channel_identity_v2", + "channel_session_v2", + "chat_message", + "chat_session", + "metadata", + "tenant", + "tool_log", + "trace_event", + "ui_session_owner", +] diff --git a/svc/persistence/ddl/postgresql_bootstrap.sql b/svc/persistence/ddl/postgresql_bootstrap.sql new file mode 100644 index 00000000..fad51181 --- /dev/null +++ b/svc/persistence/ddl/postgresql_bootstrap.sql @@ -0,0 +1,486 @@ +-- Generated for PostgreSQL assistant store (from SQLite schema) +SET client_min_messages TO WARNING; +CREATE TABLE IF NOT EXISTS chat_session ( + id TEXT PRIMARY KEY, + title TEXT NOT NULL, + created_at TEXT NOT NULL, + last_message_at TEXT + ); + +CREATE TABLE IF NOT EXISTS tenant ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + created_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS app_user ( + id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + username TEXT, + display_name TEXT NOT NULL, + role TEXT NOT NULL, + password_hash TEXT, + is_active INTEGER NOT NULL DEFAULT 1, + created_at TEXT NOT NULL, avatar_attachment_id TEXT, + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS channel_identity ( + tenant_id TEXT NOT NULL, + channel TEXT NOT NULL, + external_user_id TEXT NOT NULL, + user_id TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, channel, external_user_id), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS channel_identity_v2 ( + tenant_id TEXT NOT NULL, + channel TEXT NOT NULL, + account_id TEXT NOT NULL, + external_user_id TEXT NOT NULL, + user_id TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, channel, account_id, external_user_id), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS bind_code ( + code TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + role TEXT NOT NULL, + created_at TEXT NOT NULL, + used_at TEXT, + used_by_external_user_id TEXT, + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS channel_session ( + tenant_id TEXT NOT NULL, + channel TEXT NOT NULL, + external_chat_id TEXT NOT NULL, + external_user_id TEXT NOT NULL, + session_id TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, channel, external_chat_id, external_user_id), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(session_id) REFERENCES chat_session(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS channel_session_v2 ( + tenant_id TEXT NOT NULL, + channel TEXT NOT NULL, + account_id TEXT NOT NULL, + external_chat_id TEXT NOT NULL, + external_user_id TEXT NOT NULL, + session_id TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, channel, account_id, external_chat_id, external_user_id), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(session_id) REFERENCES chat_session(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS user_channel_account ( + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + channel TEXT NOT NULL, + account_id TEXT NOT NULL, + name TEXT NOT NULL DEFAULT '', + config TEXT NOT NULL DEFAULT '{}', + is_active INTEGER NOT NULL DEFAULT 1, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, user_id, channel, account_id), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS todo_item ( + id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + owner_user_id TEXT NOT NULL, + assignee_user_id TEXT, + title TEXT NOT NULL, + due_at TEXT, + status TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(owner_user_id) REFERENCES app_user(id) ON DELETE CASCADE, + FOREIGN KEY(assignee_user_id) REFERENCES app_user(id) ON DELETE SET NULL + ); + +CREATE TABLE IF NOT EXISTS chat_message ( + id BIGSERIAL PRIMARY KEY, + session_id TEXT NOT NULL, + role TEXT NOT NULL, + content TEXT NOT NULL, + tool_calls TEXT, + attachments TEXT, + turn_uuid TEXT, + event_type TEXT, + event_payload TEXT, + timestamp TEXT NOT NULL, + FOREIGN KEY(session_id) REFERENCES chat_session(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS ui_session_owner ( + session_id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + created_at TEXT NOT NULL, + FOREIGN KEY(session_id) REFERENCES chat_session(id) ON DELETE CASCADE, + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS user_workspace_path_allowlist ( + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + extra_roots TEXT NOT NULL DEFAULT '', + allow_any_path INTEGER NOT NULL DEFAULT 0, + allow_high_risk_public_tools INTEGER NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, user_id), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS auth_session ( + session_token_hash TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + role TEXT NOT NULL, + created_at TEXT NOT NULL, + expires_at TEXT NOT NULL, + last_seen_at TEXT NOT NULL, + revoked_at TEXT, + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS role_permission ( + role TEXT NOT NULL, + permission TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (role, permission) + ); + +CREATE TABLE IF NOT EXISTS user_permission ( + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + permission TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (tenant_id, user_id, permission), + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS admin_audit_log ( + id BIGSERIAL PRIMARY KEY, + actor_tenant_id TEXT NOT NULL, + actor_user_id TEXT NOT NULL, + action TEXT NOT NULL, + target_type TEXT NOT NULL, + target_id TEXT NOT NULL, + status TEXT NOT NULL, + detail TEXT NOT NULL DEFAULT '{}', + timestamp TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS attachment_acl ( + attachment_id TEXT NOT NULL, + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + session_id TEXT NOT NULL, + source TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (attachment_id, tenant_id, user_id, session_id, source), + FOREIGN KEY(session_id) REFERENCES chat_session(id) ON DELETE CASCADE, + FOREIGN KEY(tenant_id) REFERENCES tenant(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS tool_log ( + id BIGSERIAL PRIMARY KEY, + session_id TEXT NOT NULL, + tool_name TEXT NOT NULL, + specialist TEXT NOT NULL DEFAULT '', + args TEXT NOT NULL, + result TEXT NOT NULL, + timestamp TEXT NOT NULL, + duration_ms INTEGER, + FOREIGN KEY(session_id) REFERENCES chat_session(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS tool_plugin ( + plugin_name TEXT NOT NULL, + plugin_version TEXT NOT NULL, + entry_point TEXT NOT NULL, + enabled INTEGER NOT NULL DEFAULT 1, + updated_at TEXT NOT NULL, + PRIMARY KEY (plugin_name, entry_point) + ); + +CREATE TABLE IF NOT EXISTS mcp_server_registry ( + server_id TEXT PRIMARY KEY, + source_type TEXT NOT NULL, + source_ref TEXT NOT NULL, + version TEXT NOT NULL DEFAULT '', + entry_command TEXT NOT NULL DEFAULT '', + entry_args TEXT NOT NULL DEFAULT '[]', + env_schema TEXT NOT NULL DEFAULT '{}', + required_permissions TEXT NOT NULL DEFAULT '[]', + risk_level TEXT NOT NULL DEFAULT 'high', + timeout_s DOUBLE PRECISION NOT NULL DEFAULT 30, + enabled INTEGER NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS mcp_server_installation ( + id BIGSERIAL PRIMARY KEY, + server_id TEXT NOT NULL, + status TEXT NOT NULL, + error_code TEXT NOT NULL DEFAULT '', + detail TEXT NOT NULL DEFAULT '{}', + install_command TEXT NOT NULL DEFAULT '', + timestamp TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS mcp_server_health ( + server_id TEXT PRIMARY KEY, + status TEXT NOT NULL, + detail TEXT NOT NULL DEFAULT '{}', + checked_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS mcp_server_tool ( + server_id TEXT NOT NULL, + tool_name TEXT NOT NULL, + description TEXT NOT NULL DEFAULT '', + parameters TEXT NOT NULL DEFAULT '{}', + updated_at TEXT NOT NULL, + PRIMARY KEY (server_id, tool_name) + ); + +CREATE TABLE IF NOT EXISTS app_setting ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL, + is_secret INTEGER NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS knowledge_chunk ( + chunk_id TEXT PRIMARY KEY, + source TEXT NOT NULL, + content TEXT NOT NULL, + metadata TEXT NOT NULL DEFAULT '{}', + updated_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS knowledge_embedding ( + chunk_id TEXT NOT NULL, + model TEXT NOT NULL, + dim INTEGER NOT NULL, + vector_json TEXT NOT NULL, + updated_at TEXT NOT NULL, + PRIMARY KEY (chunk_id, model), + FOREIGN KEY(chunk_id) REFERENCES knowledge_chunk(chunk_id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS memory_item ( + memory_id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + session_id TEXT NOT NULL, + memory_type TEXT NOT NULL, + content TEXT NOT NULL, + confidence DOUBLE PRECISION NOT NULL DEFAULT 0, + source TEXT NOT NULL DEFAULT 'memory', + metadata TEXT NOT NULL DEFAULT '{}', + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + expires_at TEXT + ); + +CREATE TABLE IF NOT EXISTS memory_vector ( + memory_id TEXT NOT NULL, + model TEXT NOT NULL, + dim INTEGER NOT NULL, + vector_json TEXT NOT NULL, + updated_at TEXT NOT NULL, + PRIMARY KEY (memory_id, model), + FOREIGN KEY(memory_id) REFERENCES memory_item(memory_id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS memory_hit_log ( + id BIGSERIAL PRIMARY KEY, + tenant_id TEXT NOT NULL, + user_id TEXT NOT NULL, + session_id TEXT, + memory_id TEXT, + query_text TEXT NOT NULL, + score DOUBLE PRECISION NOT NULL DEFAULT 0, + source TEXT NOT NULL DEFAULT 'memory', + timestamp TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS oclaw_task ( + id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + session_id TEXT NOT NULL, + task_type TEXT NOT NULL DEFAULT 'async_turn', + status TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '{}', + result TEXT NOT NULL DEFAULT '{}', + attempt_count INTEGER NOT NULL DEFAULT 0, + claimed_by TEXT, + lease_expires_at TEXT, + last_error TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + finished_at TEXT + ); + +CREATE TABLE IF NOT EXISTS oclaw_run ( + run_id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + session_id TEXT NOT NULL, + status TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '{}', + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS oclaw_attempt ( + id BIGSERIAL PRIMARY KEY, + run_id TEXT NOT NULL, + tenant_id TEXT NOT NULL, + session_id TEXT NOT NULL, + attempt_no INTEGER NOT NULL, + status TEXT NOT NULL, + reason TEXT NOT NULL DEFAULT '', + payload TEXT NOT NULL DEFAULT '{}', + created_at TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS agent_audit_log ( + id BIGSERIAL PRIMARY KEY, + session_id TEXT NOT NULL, + specialist TEXT NOT NULL, + task_kind TEXT NOT NULL, + action TEXT NOT NULL, + payload TEXT NOT NULL, + status TEXT NOT NULL, + reason TEXT NOT NULL, + duration_ms INTEGER NOT NULL DEFAULT 0, + timestamp TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS agent_eval_log ( + id BIGSERIAL PRIMARY KEY, + session_id TEXT NOT NULL, + specialist TEXT NOT NULL, + task_kind TEXT NOT NULL, + success INTEGER NOT NULL, + latency_ms INTEGER NOT NULL, + cost_hint DOUBLE PRECISION NOT NULL DEFAULT 0, + notes TEXT NOT NULL DEFAULT '', + timestamp TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS trace_event ( + id BIGSERIAL PRIMARY KEY, + session_id TEXT NOT NULL, + trace_id TEXT NOT NULL, + span_id TEXT NOT NULL, + parent_span_id TEXT, + event_type TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '{}', + timestamp TEXT NOT NULL + ); + +CREATE TABLE IF NOT EXISTS llm_profile ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + mode TEXT NOT NULL, + model TEXT, + base_url TEXT, + api_key TEXT, + updated_at TEXT NOT NULL + , is_builtin INTEGER NOT NULL DEFAULT 0, hide_in_ui INTEGER NOT NULL DEFAULT 0, owner_user_id TEXT, thinking_mode_enabled INTEGER NOT NULL DEFAULT 0, reasoning_effort TEXT); + +CREATE TABLE IF NOT EXISTS llm_profile_user_grant ( + id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + profile_id TEXT NOT NULL, + user_id TEXT NOT NULL, + created_at TEXT NOT NULL, + created_by_user_id TEXT, + UNIQUE(tenant_id, profile_id, user_id), + FOREIGN KEY(profile_id) REFERENCES llm_profile(id) ON DELETE CASCADE, + FOREIGN KEY(user_id) REFERENCES app_user(id) ON DELETE CASCADE + ); + +CREATE TABLE IF NOT EXISTS llm_profile_tenant_grant ( + id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL, + profile_id TEXT NOT NULL, + created_at TEXT NOT NULL, + created_by_user_id TEXT, + UNIQUE(tenant_id, profile_id), + FOREIGN KEY(profile_id) REFERENCES llm_profile(id) ON DELETE CASCADE + ); + +CREATE INDEX IF NOT EXISTS idx_chat_session_activity ON chat_session(COALESCE(last_message_at, created_at) DESC, created_at DESC); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_app_user_tenant_username ON app_user(tenant_id, username); + +CREATE INDEX IF NOT EXISTS idx_user_channel_account_channel_account ON user_channel_account(channel, account_id, is_active); + +CREATE INDEX IF NOT EXISTS idx_chat_message_session_turn_uuid ON chat_message(session_id, turn_uuid); + +CREATE INDEX IF NOT EXISTS idx_ui_session_owner_tenant_user ON ui_session_owner(tenant_id, user_id, created_at DESC); + +CREATE INDEX IF NOT EXISTS idx_ui_session_owner_tenant_user_session ON ui_session_owner(tenant_id, user_id, session_id); + +CREATE INDEX IF NOT EXISTS idx_ui_session_owner_tenant_session ON ui_session_owner(tenant_id, session_id); + +CREATE INDEX IF NOT EXISTS idx_auth_session_user_expires ON auth_session(user_id, expires_at); + +CREATE INDEX IF NOT EXISTS idx_admin_audit_actor_ts ON admin_audit_log(actor_user_id, timestamp DESC); + +CREATE INDEX IF NOT EXISTS idx_attachment_acl_tenant_attachment ON attachment_acl(tenant_id, attachment_id, created_at DESC); + +CREATE INDEX IF NOT EXISTS idx_attachment_acl_user_attachment ON attachment_acl(tenant_id, user_id, attachment_id, created_at DESC); + +CREATE INDEX IF NOT EXISTS idx_attachment_acl_session_attachment ON attachment_acl(session_id, attachment_id, created_at DESC); + +CREATE INDEX IF NOT EXISTS idx_memory_item_tenant_user_updated ON memory_item(tenant_id, user_id, updated_at DESC); + +CREATE INDEX IF NOT EXISTS idx_memory_hit_log_tenant_user_ts ON memory_hit_log(tenant_id, user_id, timestamp DESC); + +CREATE INDEX IF NOT EXISTS idx_memory_item_session_updated ON memory_item(session_id, updated_at DESC); + +CREATE INDEX IF NOT EXISTS idx_oclaw_task_status_updated ON oclaw_task(status, updated_at); + +CREATE INDEX IF NOT EXISTS idx_oclaw_task_tenant_session ON oclaw_task(tenant_id, session_id, created_at DESC); + +CREATE INDEX IF NOT EXISTS idx_oclaw_run_tenant_session ON oclaw_run(tenant_id, session_id, updated_at DESC); + +CREATE INDEX IF NOT EXISTS idx_oclaw_attempt_run_no ON oclaw_attempt(run_id, attempt_no); + +CREATE INDEX IF NOT EXISTS idx_knowledge_chunk_source_updated ON knowledge_chunk(source, updated_at); + +CREATE INDEX IF NOT EXISTS idx_trace_event_session_id_id ON trace_event(session_id, id); + +CREATE INDEX IF NOT EXISTS idx_llm_profile_grant_user ON llm_profile_user_grant(tenant_id, user_id); + +CREATE INDEX IF NOT EXISTS idx_llm_profile_grant_profile ON llm_profile_user_grant(tenant_id, profile_id); + +CREATE INDEX IF NOT EXISTS idx_llm_profile_tenant_grant ON llm_profile_tenant_grant(tenant_id, profile_id); + +CREATE INDEX IF NOT EXISTS idx_chat_message_session_id_id ON chat_message(session_id, id); diff --git a/svc/persistence/pg_adapter.py b/svc/persistence/pg_adapter.py new file mode 100644 index 00000000..9583d6fa --- /dev/null +++ b/svc/persistence/pg_adapter.py @@ -0,0 +1,68 @@ +"""psycopg connection surface compatible with sqlite3 usage in SqliteStore.""" + +from __future__ import annotations + +from typing import Any, Iterable, Sequence + +import psycopg +from psycopg.rows import dict_row + +from svc.persistence.pg_compat import adapt_sql_for_postgres + + +def normalize_psycopg_conninfo(url: str) -> str: + """Strip SQLAlchemy driver suffix so :func:`psycopg.connect` accepts the URI.""" + u = str(url or "").strip() + for prefix in ( + "postgresql+psycopg://", + "postgresql+psycopg2://", + "postgres+psycopg://", + "postgres+psycopg2://", + ): + if u.startswith(prefix): + rest = u.split("://", 1)[1] + return "postgresql://" + rest + return u + + +class PgCursorShim: + def __init__(self, raw: Any) -> None: + self._raw = raw + + def fetchone(self) -> Any: + return self._raw.fetchone() + + def fetchall(self) -> list[Any]: + return self._raw.fetchall() + + @property + def lastrowid(self) -> int: + return 0 + + @property + def rowcount(self) -> int: + return int(self._raw.rowcount or 0) + + def __iter__(self) -> Iterable[Any]: + return iter(self._raw) + + +class PgConnShim: + def __init__(self, raw: psycopg.Connection) -> None: + self._raw = raw + + def execute(self, sql: str, params: Sequence[Any] | None = None) -> PgCursorShim: + adapted = adapt_sql_for_postgres(sql) + cur = self._raw.execute(adapted, params or ()) + return PgCursorShim(cur) + + def executemany(self, sql: str, seq_of_params: Sequence[Sequence[Any]]) -> None: + adapted = adapt_sql_for_postgres(sql) + self._raw.executemany(adapted, seq_of_params) + + +def connect_postgres(url: str) -> psycopg.Connection: + return psycopg.connect(normalize_psycopg_conninfo(url), row_factory=dict_row) + + +__all__ = ["PgConnShim", "PgCursorShim", "connect_postgres", "normalize_psycopg_conninfo"] diff --git a/svc/persistence/pg_compat.py b/svc/persistence/pg_compat.py new file mode 100644 index 00000000..a3bbe8d5 --- /dev/null +++ b/svc/persistence/pg_compat.py @@ -0,0 +1,230 @@ +"""Translate SQLite-oriented SQL to PostgreSQL for psycopg execution.""" + +from __future__ import annotations + +import re +from typing import Any + + +def scrub_nul_bytes_from_text(s: str | None) -> str | None: + """PostgreSQL ``TEXT`` / ``VARCHAR`` reject U+0000; SQLite allows it. + + Strip NULs from any string bound for PG text columns so assistant/tool rows + persist instead of failing the whole ``INSERT`` after the user row succeeded. + """ + if s is None: + return None + if "\x00" not in s: + return s + return s.replace("\x00", "") + + +def scrub_nul_bytes_from_jsonable(obj: Any) -> Any: + """Recursively remove NUL from strings inside dict/list before ``json.dumps``. + + ``json.dumps`` encodes embedded NUL as the six-character ``\\u0000`` escape; a + plain ``TEXT`` scrub on the serialized JSON would not remove the decoded NUL + after reload, and PostgreSQL still rejects a true NUL inside string values. + """ + if isinstance(obj, str): + return obj.replace("\x00", "") if "\x00" in obj else obj + if isinstance(obj, dict): + return {k: scrub_nul_bytes_from_jsonable(v) for k, v in obj.items()} + if isinstance(obj, list): + return [scrub_nul_bytes_from_jsonable(v) for v in obj] + if isinstance(obj, tuple): + return tuple(scrub_nul_bytes_from_jsonable(v) for v in obj) + return obj + + +def qmarks_to_percent(sql: str) -> str: + """Replace ``?`` placeholders outside single-quoted strings with ``%s`` (psycopg).""" + out: list[str] = [] + i = 0 + in_single = False + while i < len(sql): + ch = sql[i] + if ch == "'" and (i == 0 or sql[i - 1] != "\\"): + in_single = not in_single + out.append(ch) + i += 1 + continue + if ch == "?" and not in_single: + out.append("%s") + else: + out.append(ch) + i += 1 + return "".join(out) + + +def rewrite_sqlite_extensions_for_postgres(sql: str) -> str: + """Rewrite SQLite-only INSERT forms to PostgreSQL-compatible SQL (still uses ``?``).""" + s = sql + repls: list[tuple[str, str]] = [ + ( + """INSERT OR REPLACE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + VALUES (?, ?, ?, ?)""", + """INSERT INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + VALUES (?, ?, ?, ?) + ON CONFLICT (session_id) DO UPDATE SET + tenant_id = EXCLUDED.tenant_id, + user_id = EXCLUDED.user_id, + created_at = EXCLUDED.created_at""", + ), + ( + """INSERT OR IGNORE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + SELECT DISTINCT cs.session_id, cs.tenant_id, ci.user_id, ? + FROM channel_session_v2 cs + JOIN channel_identity_v2 ci + ON ci.tenant_id = cs.tenant_id + AND ci.channel = cs.channel + AND ci.account_id = cs.account_id + AND ci.external_user_id = cs.external_user_id + WHERE cs.session_id IS NOT NULL AND cs.session_id != ''""", + """INSERT INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + SELECT DISTINCT cs.session_id, cs.tenant_id, ci.user_id, ? + FROM channel_session_v2 cs + JOIN channel_identity_v2 ci + ON ci.tenant_id = cs.tenant_id + AND ci.channel = cs.channel + AND ci.account_id = cs.account_id + AND ci.external_user_id = cs.external_user_id + WHERE cs.session_id IS NOT NULL AND cs.session_id != '' + ON CONFLICT (session_id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + SELECT s.id, ?, ?, COALESCE(s.created_at, ?) + FROM chat_session s + WHERE NOT EXISTS (SELECT 1 FROM ui_session_owner o WHERE o.session_id = s.id)""", + """INSERT INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + SELECT s.id, ?, ?, COALESCE(s.created_at, ?) + FROM chat_session s + WHERE NOT EXISTS (SELECT 1 FROM ui_session_owner o WHERE o.session_id = s.id) + ON CONFLICT (session_id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + VALUES (?, ?, ?, ?)""", + """INSERT INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + VALUES (?, ?, ?, ?) + ON CONFLICT (session_id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO llm_profile + (id, name, mode, model, base_url, api_key, updated_at, is_builtin, hide_in_ui, owner_user_id) + VALUES (?, ?, 'ollama', ?, ?, NULL, ?, 1, 0, NULL)""", + """INSERT INTO llm_profile + (id, name, mode, model, base_url, api_key, updated_at, is_builtin, hide_in_ui, owner_user_id) + VALUES (?, ?, 'ollama', ?, ?, NULL, ?, 1, 0, NULL) + ON CONFLICT (id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO llm_profile + (id, name, mode, model, base_url, api_key, updated_at, is_builtin, hide_in_ui, owner_user_id) + VALUES (?, ?, 'rule', NULL, NULL, NULL, ?, 1, 1, NULL)""", + """INSERT INTO llm_profile + (id, name, mode, model, base_url, api_key, updated_at, is_builtin, hide_in_ui, owner_user_id) + VALUES (?, ?, 'rule', NULL, NULL, NULL, ?, 1, 1, NULL) + ON CONFLICT (id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO role_permission(role, permission, created_at) + VALUES (?, ?, ?)""", + """INSERT INTO role_permission(role, permission, created_at) + VALUES (?, ?, ?) + ON CONFLICT (role, permission) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO attachment_acl + (attachment_id, tenant_id, user_id, session_id, source, created_at) + VALUES (?, ?, ?, ?, ?, ?)""", + """INSERT INTO attachment_acl + (attachment_id, tenant_id, user_id, session_id, source, created_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT (attachment_id, tenant_id, user_id, session_id, source) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO attachment_acl + (attachment_id, tenant_id, user_id, session_id, source, created_at) + VALUES (?, ?, ?, ?, ?, ?)""", + """INSERT INTO attachment_acl + (attachment_id, tenant_id, user_id, session_id, source, created_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT (attachment_id, tenant_id, user_id, session_id, source) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO llm_profile_user_grant + (id, tenant_id, profile_id, user_id, created_at, created_by_user_id) + VALUES (?, ?, ?, ?, ?, ?)""", + """INSERT INTO llm_profile_user_grant + (id, tenant_id, profile_id, user_id, created_at, created_by_user_id) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT (tenant_id, profile_id, user_id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO llm_profile_tenant_grant + (id, tenant_id, profile_id, created_at, created_by_user_id) + VALUES (?, ?, ?, ?, ?)""", + """INSERT INTO llm_profile_tenant_grant + (id, tenant_id, profile_id, created_at, created_by_user_id) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT (tenant_id, profile_id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO user_permission (tenant_id, user_id, permission, created_at) + VALUES (?, ?, ?, ?)""", + """INSERT INTO user_permission (tenant_id, user_id, permission, created_at) + VALUES (?, ?, ?, ?) + ON CONFLICT (tenant_id, user_id, permission) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO channel_session + (tenant_id, channel, external_chat_id, external_user_id, session_id, created_at) + VALUES (?, ?, ?, ?, ?, ?)""", + """INSERT INTO channel_session + (tenant_id, channel, external_chat_id, external_user_id, session_id, created_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT (tenant_id, channel, external_chat_id, external_user_id) DO NOTHING""", + ), + ( + """INSERT OR IGNORE INTO channel_session_v2 + (tenant_id, channel, account_id, external_chat_id, external_user_id, session_id, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?)""", + """INSERT INTO channel_session_v2 + (tenant_id, channel, account_id, external_chat_id, external_user_id, session_id, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?) + ON CONFLICT (tenant_id, channel, account_id, external_chat_id, external_user_id) DO NOTHING""", + ), + ] + for old, new in repls: + if old in s: + s = s.replace(old, new, 1) + if "INSERT OR IGNORE INTO" in s: + s = re.sub( + r"INSERT\s+OR\s+IGNORE\s+INTO\s+(\w+)\s+", + r"INSERT INTO \1 ", + s, + count=1, + flags=re.IGNORECASE | re.DOTALL, + ) + if "ON CONFLICT" not in s.upper(): + s = s.rstrip() + "\nON CONFLICT DO NOTHING" + if "INSERT OR REPLACE INTO" in s.upper(): + raise ValueError( + "unsupported INSERT OR REPLACE for PostgreSQL; extend svc.persistence.pg_compat" + ) + return s + + +def adapt_sql_for_postgres(sql: str) -> str: + return qmarks_to_percent(rewrite_sqlite_extensions_for_postgres(sql)) + + +__all__ = [ + "adapt_sql_for_postgres", + "qmarks_to_percent", + "rewrite_sqlite_extensions_for_postgres", + "scrub_nul_bytes_from_jsonable", + "scrub_nul_bytes_from_text", +] diff --git a/svc/persistence/sa_repos/__init__.py b/svc/persistence/sa_repos/__init__.py new file mode 100644 index 00000000..3d209c34 --- /dev/null +++ b/svc/persistence/sa_repos/__init__.py @@ -0,0 +1,30 @@ +"""SQLAlchemy-backed repository slices (incremental migration off raw SQL).""" + +from __future__ import annotations + +from svc.persistence.sa_repos.admin_user_stats import AdminUserStatsSaRepository +from svc.persistence.sa_repos.app_settings import AppSettingsSaRepository +from svc.persistence.sa_repos.app_users import AppUsersSaRepository +from svc.persistence.sa_repos.auth_sessions import AuthSessionsSaRepository +from svc.persistence.sa_repos.chat_messages import ChatMessagesSaRepository +from svc.persistence.sa_repos.chat_sessions import ChatSessionsSaRepository +from svc.persistence.sa_repos.session_tool_health import SessionToolHealthSaRepository +from svc.persistence.sa_repos.tenant_bind_code import BindCodeSaRepository, TenantSaRepository +from svc.persistence.sa_repos.tool_log_queries import ToolLogQueriesSaRepository +from svc.persistence.sa_repos.trace_events import TraceEventsSaRepository +from svc.persistence.sa_repos.ui_session_owner import UiSessionOwnerSaRepository + +__all__ = [ + "AdminUserStatsSaRepository", + "AppSettingsSaRepository", + "AppUsersSaRepository", + "AuthSessionsSaRepository", + "BindCodeSaRepository", + "ChatMessagesSaRepository", + "ChatSessionsSaRepository", + "SessionToolHealthSaRepository", + "TenantSaRepository", + "ToolLogQueriesSaRepository", + "TraceEventsSaRepository", + "UiSessionOwnerSaRepository", +] diff --git a/svc/persistence/sa_repos/admin_user_stats.py b/svc/persistence/sa_repos/admin_user_stats.py new file mode 100644 index 00000000..faffde76 --- /dev/null +++ b/svc/persistence/sa_repos/admin_user_stats.py @@ -0,0 +1,167 @@ +"""Admin tenant user stats (list_admin_user_stats) via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy import and_, case, distinct, func, literal, or_, select +from sqlalchemy.engine import Engine +from sqlalchemy.sql import bindparam + +from svc.persistence.db.tables import ( + app_user, + auth_session, + chat_session, + trace_event, + ui_session_owner, +) + + +class AdminUserStatsSaRepository: + """Aggregates for ``SqliteStore.list_admin_user_stats``.""" + + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def fetch( + self, + *, + tenant_id: str, + search_lower: str | None, + cutoff_iso: str, + limit: int, + offset: int, + ) -> dict[str, Any]: + tid = str(tenant_id or "").strip() + lim = max(1, min(int(limit), 500)) + off = max(0, int(offset)) + q_text = str(search_lower or "").strip().lower() or None + cutoff = str(cutoff_iso) + + conds: list[Any] = [app_user.c.tenant_id == tid] + if q_text: + like = f"%{q_text}%" + conds.append( + or_( + func.lower(func.coalesce(app_user.c.username, literal(""))).like(like), + func.lower(func.coalesce(app_user.c.display_name, literal(""))).like(like), + ) + ) + wh = and_(*conds) + + user_stmt = ( + select( + app_user.c.id.label("user_id"), + app_user.c.username, + func.coalesce(app_user.c.display_name, literal("")).label("display_name"), + app_user.c.role, + app_user.c.is_active, + ) + .where(wh) + .order_by(app_user.c.username.asc()) + .limit(lim) + .offset(off) + ) + cnt_stmt = select(func.count()).select_from(app_user).where(wh) + + sess_join = chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + active_ts = func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at) + total_active_sess_stmt = ( + select(func.count(distinct(chat_session.c.id))) + .select_from(sess_join) + .where(ui_session_owner.c.tenant_id == tid, active_ts >= literal(cutoff)) + ) + total_active_logins_stmt = ( + select(func.count()) + .select_from(auth_session) + .where( + auth_session.c.tenant_id == tid, + auth_session.c.revoked_at.is_(None), + auth_session.c.expires_at > literal(cutoff), + auth_session.c.last_seen_at >= literal(cutoff), + ) + ) + + with self._engine.connect() as conn: + total_users = int(conn.execute(cnt_stmt).scalar_one() or 0) + user_rows = [dict(r) for r in conn.execute(user_stmt).mappings().all()] + total_active_sessions = int(conn.execute(total_active_sess_stmt).scalar_one() or 0) + total_active_logins = int(conn.execute(total_active_logins_stmt).scalar_one() or 0) + + uids = [str(r["user_id"] or "").strip() for r in user_rows if str(r.get("user_id") or "").strip()] + trace_rows: list[dict[str, Any]] = [] + own_count_rows: list[dict[str, Any]] = [] + active_sess_rows: list[dict[str, Any]] = [] + login_rows: list[dict[str, Any]] = [] + + if uids: + uids_param = bindparam("uids", expanding=True) + trace_stmt = ( + select(ui_session_owner.c.user_id, trace_event.c.payload) + .select_from( + trace_event.join( + ui_session_owner, + ui_session_owner.c.session_id == trace_event.c.session_id, + ) + ) + .where(ui_session_owner.c.tenant_id == tid, ui_session_owner.c.user_id.in_(uids_param)) + ) + own_cnt_stmt = ( + select(ui_session_owner.c.user_id, func.count().label("c")) + .where(ui_session_owner.c.tenant_id == tid, ui_session_owner.c.user_id.in_(uids_param)) + .group_by(ui_session_owner.c.user_id) + ) + active_case = case( + (active_ts >= literal(cutoff), chat_session.c.id), + else_=None, + ) + active_sess_stmt = ( + select( + ui_session_owner.c.user_id, + func.count(distinct(active_case)).label("active_30m"), + func.max(active_ts).label("last_message_at"), + ) + .select_from(sess_join) + .where(ui_session_owner.c.tenant_id == tid, ui_session_owner.c.user_id.in_(uids_param)) + .group_by(ui_session_owner.c.user_id) + ) + login_stmt = ( + select( + auth_session.c.user_id, + func.count().label("c"), + func.max(auth_session.c.last_seen_at).label("last_seen_at"), + ) + .where( + auth_session.c.tenant_id == tid, + auth_session.c.user_id.in_(uids_param), + auth_session.c.revoked_at.is_(None), + auth_session.c.expires_at > literal(cutoff), + auth_session.c.last_seen_at >= literal(cutoff), + ) + .group_by(auth_session.c.user_id) + ) + bind = {"uids": uids} + with self._engine.connect() as conn: + trace_rows = [dict(r) for r in conn.execute(trace_stmt, bind).mappings().all()] + own_count_rows = [dict(r) for r in conn.execute(own_cnt_stmt, bind).mappings().all()] + active_sess_rows = [dict(r) for r in conn.execute(active_sess_stmt, bind).mappings().all()] + login_rows = [dict(r) for r in conn.execute(login_stmt, bind).mappings().all()] + + return { + "total_users": total_users, + "user_rows": user_rows, + "total_active_sessions_30m": total_active_sessions, + "total_active_logins_30m": total_active_logins, + "trace_rows": trace_rows, + "sessions_count_rows": own_count_rows, + "active_sess_rows": active_sess_rows, + "login_rows": login_rows, + } + + +__all__ = ["AdminUserStatsSaRepository"] diff --git a/svc/persistence/sa_repos/app_settings.py b/svc/persistence/sa_repos/app_settings.py new file mode 100644 index 00000000..4b755c94 --- /dev/null +++ b/svc/persistence/sa_repos/app_settings.py @@ -0,0 +1,141 @@ +"""app_setting access via SQLAlchemy Core (SQLite + PostgreSQL).""" + +from __future__ import annotations + +from collections.abc import Callable + +from sqlalchemy import delete, func, select, update +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.dialects.sqlite import insert as sqlite_insert +from sqlalchemy.engine import Connection, Engine + +from svc.persistence.db.tables import app_setting + + +class AppSettingsSaRepository: + """Phase-1 SA migration: ``app_setting`` reads/writes only.""" + + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def _dialect(self, conn: Connection) -> str: + return conn.engine.dialect.name + + def _upsert( + self, + conn: Connection, + *, + key: str, + value: str, + is_secret: int, + updated_at: str, + ) -> None: + dialect = self._dialect(conn) + if dialect == "sqlite": + ins = sqlite_insert(app_setting).values( + key=key, + value=value, + is_secret=is_secret, + updated_at=updated_at, + ) + stmt = ins.on_conflict_do_update( + index_elements=[app_setting.c.key], + set_={ + "value": ins.excluded.value, + "is_secret": ins.excluded.is_secret, + "updated_at": ins.excluded.updated_at, + }, + ) + elif dialect == "postgresql": + ins = pg_insert(app_setting).values( + key=key, + value=value, + is_secret=is_secret, + updated_at=updated_at, + ) + stmt = ins.on_conflict_do_update( + index_elements=[app_setting.c.key], + set_={ + "value": ins.excluded.value, + "is_secret": ins.excluded.is_secret, + "updated_at": ins.excluded.updated_at, + }, + ) + else: + raise RuntimeError(f"unsupported SQLAlchemy dialect for app_setting: {dialect!r}") + conn.execute(stmt) + + def upsert_plain(self, *, key: str, value: str, updated_at: str) -> None: + with self._engine.begin() as conn: + self._upsert(conn, key=key, value=value, is_secret=0, updated_at=updated_at) + + def upsert_secret(self, *, key: str, encoded_value: str, updated_at: str) -> None: + with self._engine.begin() as conn: + self._upsert(conn, key=key, value=encoded_value, is_secret=1, updated_at=updated_at) + + def fetch_row(self, *, key: str) -> tuple[str, int] | None: + with self._engine.connect() as conn: + row = conn.execute( + select(app_setting.c.value, app_setting.c.is_secret).where(app_setting.c.key == key) + ).one_or_none() + if row is None: + return None + return (str(row[0]), int(row[1])) + + def delete_key(self, *, key: str) -> None: + with self._engine.begin() as conn: + conn.execute(delete(app_setting).where(app_setting.c.key == key)) + + def migrate_b64_secrets( + self, + *, + ts: str, + decode_secret: Callable[[str], str], + encode_secret: Callable[[str], str], + predicate_new_encoding: Callable[[str], bool], + ) -> int: + """Re-encode legacy ``b64:`` rows; idempotent. Returns rows updated.""" + migrated = 0 + with self._engine.begin() as conn: + rows = conn.execute( + select(app_setting.c.key, app_setting.c.value).where( + app_setting.c.is_secret == 1, + app_setting.c.value.like("b64:%"), + ) + ).all() + for k, v in rows: + key = str(k or "") + val = str(v or "") + if not key: + continue + try: + plain = decode_secret(val) + except Exception: + continue + enc = encode_secret(plain) + if enc != val and predicate_new_encoding(enc): + conn.execute( + update(app_setting) + .where( + app_setting.c.key == key, + app_setting.c.is_secret == 1, + ) + .values(value=enc, updated_at=ts) + ) + migrated += 1 + return migrated + + def count_legacy_b64_secrets(self) -> int: + with self._engine.connect() as conn: + n = conn.execute( + select(func.count()).select_from(app_setting).where( + app_setting.c.is_secret == 1, + app_setting.c.value.like("b64:%"), + ) + ).scalar_one() + return int(n or 0) + + +__all__ = ["AppSettingsSaRepository"] diff --git a/svc/persistence/sa_repos/app_users.py b/svc/persistence/sa_repos/app_users.py new file mode 100644 index 00000000..84cc3562 --- /dev/null +++ b/svc/persistence/sa_repos/app_users.py @@ -0,0 +1,296 @@ +"""app_user CRUD slices via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any, Mapping + +from sqlalchemy import case, delete, exists, func, insert, literal, or_, select, union, update +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import app_user, channel_identity, channel_identity_v2 + + +def _user_row_to_public_dict(r: Mapping[str, Any]) -> dict[str, Any]: + return { + "id": r["id"], + "tenant_id": r["tenant_id"], + "username": r["username"], + "display_name": r["display_name"], + "role": r["role"], + "is_active": bool(int(r["is_active"] or 0)), + "created_at": r["created_at"], + "password_hash": r["password_hash"], + "avatar_attachment_id": str(r["avatar_attachment_id"] or "").strip() or None, + } + + +class AppUsersSaRepository: + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def count_by_tenant_username(self, *, tenant_id: str, username: str) -> int: + tid, un = str(tenant_id), str(username) + with self._engine.connect() as conn: + n = conn.execute( + select(func.count()) + .select_from(app_user) + .where(app_user.c.tenant_id == tid, app_user.c.username == un) + ).scalar_one() + return int(n or 0) + + def insert_user( + self, + *, + user_id: str, + tenant_id: str, + username: str, + display_name: str, + role: str, + password_hash: str, + is_active: int, + created_at: str, + avatar_attachment_id: str | None = None, + ) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(app_user).values( + id=str(user_id), + tenant_id=str(tenant_id), + username=str(username), + display_name=str(display_name), + role=str(role), + password_hash=str(password_hash or ""), + is_active=int(is_active), + created_at=str(created_at), + avatar_attachment_id=avatar_attachment_id, + ) + ) + + def _select_user_columns(self): + return select( + app_user.c.id, + app_user.c.tenant_id, + app_user.c.username, + app_user.c.display_name, + app_user.c.role, + func.coalesce(app_user.c.is_active, literal(1)).label("is_active"), + app_user.c.created_at, + func.coalesce(app_user.c.password_hash, literal("")).label("password_hash"), + func.coalesce(app_user.c.avatar_attachment_id, literal("")).label("avatar_attachment_id"), + ) + + def fetch_by_tenant_and_id(self, *, tenant_id: str, user_id: str) -> dict[str, Any] | None: + tid, uid = str(tenant_id), str(user_id) + stmt = self._select_user_columns().where(app_user.c.tenant_id == tid, app_user.c.id == uid).limit(1) + with self._engine.connect() as conn: + row = conn.execute(stmt).mappings().first() + return _user_row_to_public_dict(row) if row else None + + def fetch_by_tenant_and_username(self, *, tenant_id: str, username: str) -> dict[str, Any] | None: + tid, un = str(tenant_id), str(username) + stmt = ( + self._select_user_columns() + .where(app_user.c.tenant_id == tid, app_user.c.username == un) + .limit(1) + ) + with self._engine.connect() as conn: + row = conn.execute(stmt).mappings().first() + return _user_row_to_public_dict(row) if row else None + + def fetch_first_by_username_global(self, *, username: str) -> dict[str, Any] | None: + un = str(username) + stmt = ( + self._select_user_columns() + .where(app_user.c.username == un) + .order_by(app_user.c.created_at.asc()) + .limit(1) + ) + with self._engine.connect() as conn: + row = conn.execute(stmt).mappings().first() + return _user_row_to_public_dict(row) if row else None + + def list_users_for_tenant( + self, + *, + tenant_id: str, + limit: int, + offset: int, + q: str | None, + include_inactive: bool, + ) -> list[dict[str, Any]]: + tid = str(tenant_id) + lim = max(1, int(limit)) + off = max(0, int(offset)) + + has_password = case( + (func.trim(func.coalesce(app_user.c.password_hash, literal(""))) != literal(""), 1), + else_=0, + ).label("has_password") + + wecom_linked = or_( + exists( + select(literal(1)) + .select_from(channel_identity_v2) + .where( + channel_identity_v2.c.tenant_id == app_user.c.tenant_id, + channel_identity_v2.c.user_id == app_user.c.id, + channel_identity_v2.c.channel == literal("wecom"), + ) + ), + exists( + select(literal(1)) + .select_from(channel_identity) + .where( + channel_identity.c.tenant_id == app_user.c.tenant_id, + channel_identity.c.user_id == app_user.c.id, + channel_identity.c.channel == literal("wecom"), + ) + ), + ).label("wecom_linked") + + channel_linked = or_( + exists( + select(literal(1)) + .select_from(channel_identity_v2) + .where( + channel_identity_v2.c.tenant_id == app_user.c.tenant_id, + channel_identity_v2.c.user_id == app_user.c.id, + ) + ), + exists( + select(literal(1)) + .select_from(channel_identity) + .where( + channel_identity.c.tenant_id == app_user.c.tenant_id, + channel_identity.c.user_id == app_user.c.id, + ) + ), + ).label("channel_linked") + + eid_ci = func.trim(func.coalesce(channel_identity.c.external_user_id, literal(""))) + sq1 = ( + select(eid_ci.label("eid")) + .where( + channel_identity.c.tenant_id == app_user.c.tenant_id, + channel_identity.c.user_id == app_user.c.id, + channel_identity.c.channel == literal("wecom"), + eid_ci != literal(""), + ) + .distinct() + ) + eid_v2 = func.trim(func.coalesce(channel_identity_v2.c.external_user_id, literal(""))) + sq2 = ( + select(eid_v2.label("eid")) + .where( + channel_identity_v2.c.tenant_id == app_user.c.tenant_id, + channel_identity_v2.c.user_id == app_user.c.id, + channel_identity_v2.c.channel == literal("wecom"), + eid_v2 != literal(""), + ) + .distinct() + ) + u_sub = union(sq1, sq2).subquery() + if self._engine.dialect.name == "postgresql": + wecom_ids_expr = select(func.string_agg(u_sub.c.eid, literal(", "))).scalar_subquery() + else: + wecom_ids_expr = select(func.group_concat(u_sub.c.eid, literal(", "))).scalar_subquery() + + stmt = ( + select( + app_user.c.id, + app_user.c.tenant_id, + app_user.c.username, + app_user.c.display_name, + app_user.c.role, + func.coalesce(app_user.c.is_active, literal(1)).label("is_active"), + app_user.c.created_at, + has_password, + wecom_linked, + channel_linked, + wecom_ids_expr.label("wecom_external_user_ids"), + ) + .where(app_user.c.tenant_id == tid) + ) + token = str(q or "").strip() + if token: + key = f"%{token.lower()}%" + stmt = stmt.where( + or_( + func.lower(app_user.c.display_name).like(key), + func.lower(func.coalesce(app_user.c.username, literal(""))).like(key), + app_user.c.id.like(f"%{token[:32]}%"), + ) + ) + if not include_inactive: + stmt = stmt.where(func.coalesce(app_user.c.is_active, literal(1)) == literal(1)) + stmt = stmt.order_by(app_user.c.created_at.desc()).limit(lim).offset(off) + + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + out: list[dict[str, Any]] = [] + for r in rows: + has_pw = bool(int(r["has_password"] or 0)) + uname = str(r["username"] or "") + can_chat = bool(has_pw) + wl = r["wecom_linked"] + cl = r["channel_linked"] + out.append( + { + "id": r["id"], + "tenant_id": r["tenant_id"], + "username": r["username"], + "display_name": r["display_name"], + "role": r["role"], + "is_active": bool(int(r["is_active"] or 0)), + "created_at": r["created_at"], + "has_password": has_pw, + "wecom_linked": bool(int(wl or 0)), + "channel_linked": bool(int(cl or 0)), + "can_chat_login": can_chat, + "wecom_external_user_ids": str(r["wecom_external_user_ids"] or "").strip(), + } + ) + return out + + def update_user_account( + self, + *, + tenant_id: str, + user_id: str, + display_name: str | None = None, + role: str | None = None, + is_active: bool | None = None, + password_hash: str | None = None, + avatar_attachment_id: str | None = None, + ) -> bool: + tid, uid = str(tenant_id), str(user_id) + vals: dict[str, Any] = {} + if display_name is not None: + vals["display_name"] = str(display_name).strip() or "User" + if role is not None: + vals["role"] = str(role).strip() or "member" + if is_active is not None: + vals["is_active"] = 1 if is_active else 0 + if password_hash is not None: + vals["password_hash"] = str(password_hash) + if avatar_attachment_id is not None: + aid = str(avatar_attachment_id).strip() + vals["avatar_attachment_id"] = aid if aid else None + if not vals: + return False + with self._engine.begin() as conn: + res = conn.execute( + update(app_user).where(app_user.c.tenant_id == tid, app_user.c.id == uid).values(**vals) + ) + return bool(int(res.rowcount or 0) > 0) + + def delete_user_account(self, *, tenant_id: str, user_id: str) -> int: + tid, uid = str(tenant_id), str(user_id) + with self._engine.begin() as conn: + res = conn.execute(delete(app_user).where(app_user.c.tenant_id == tid, app_user.c.id == uid)) + return int(res.rowcount or 0) + + +__all__ = ["AppUsersSaRepository"] diff --git a/svc/persistence/sa_repos/auth_sessions.py b/svc/persistence/sa_repos/auth_sessions.py new file mode 100644 index 00000000..5dccf754 --- /dev/null +++ b/svc/persistence/sa_repos/auth_sessions.py @@ -0,0 +1,102 @@ +"""auth_session access via SQLAlchemy Core (SQLite + PostgreSQL).""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy import insert, select, update +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import auth_session + + +class AuthSessionsSaRepository: + """Phase-2 SA migration: admin login session rows.""" + + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_session( + self, + *, + session_token_hash: str, + tenant_id: str, + user_id: str, + role: str, + created_at: str, + expires_at: str, + last_seen_at: str, + ) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(auth_session).values( + session_token_hash=str(session_token_hash), + tenant_id=str(tenant_id), + user_id=str(user_id), + role=str(role), + created_at=str(created_at), + expires_at=str(expires_at), + last_seen_at=str(last_seen_at), + revoked_at=None, + ) + ) + + def revoke_one(self, *, session_token_hash: str, revoked_at: str) -> int: + with self._engine.begin() as conn: + res = conn.execute( + update(auth_session) + .where( + auth_session.c.session_token_hash == str(session_token_hash), + auth_session.c.revoked_at.is_(None), + ) + .values(revoked_at=str(revoked_at)) + ) + n = res.rowcount + if n is None or n < 0: + return 0 + return int(n) + + def revoke_all_active(self, *, revoked_at: str) -> int: + with self._engine.begin() as conn: + res = conn.execute( + update(auth_session) + .where(auth_session.c.revoked_at.is_(None)) + .values(revoked_at=str(revoked_at)) + ) + n = res.rowcount + if n is None or n < 0: + return 0 + return int(n) + + def fetch_by_hash(self, *, session_token_hash: str) -> dict[str, Any] | None: + with self._engine.connect() as conn: + row = conn.execute( + select( + auth_session.c.session_token_hash, + auth_session.c.tenant_id, + auth_session.c.user_id, + auth_session.c.role, + auth_session.c.created_at, + auth_session.c.expires_at, + auth_session.c.last_seen_at, + auth_session.c.revoked_at, + ) + .where(auth_session.c.session_token_hash == str(session_token_hash)) + .limit(1) + ).mappings().first() + if row is None: + return None + return dict(row) + + def touch(self, *, session_token_hash: str, last_seen_at: str) -> None: + with self._engine.begin() as conn: + conn.execute( + update(auth_session) + .where(auth_session.c.session_token_hash == str(session_token_hash)) + .values(last_seen_at=str(last_seen_at)) + ) + + +__all__ = ["AuthSessionsSaRepository"] diff --git a/svc/persistence/sa_repos/chat_messages.py b/svc/persistence/sa_repos/chat_messages.py new file mode 100644 index 00000000..825bc447 --- /dev/null +++ b/svc/persistence/sa_repos/chat_messages.py @@ -0,0 +1,481 @@ +"""chat_message access via SQLAlchemy Core.""" + +from __future__ import annotations + +import json +from typing import Any, Mapping + +from sqlalchemy import delete, exists, func, insert, literal, select, update +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import chat_message, chat_session +from svc.persistence.sqlite_store import ( + ChatMessage, + SessionMessagesMeta, + _tool_row_assistant_message_id, + _trim_messages_start_index, + utc_now_iso, +) + + +def _sql_text_required(v: Any) -> str: + if v is None: + return "" + if isinstance(v, (bytes, bytearray, memoryview)): + try: + return bytes(v).decode("utf-8", errors="replace") + except Exception: + return "" + return str(v) + + +def _sql_text_optional_plain(v: Any) -> str | None: + if v is None: + return None + if isinstance(v, (bytes, bytearray, memoryview)): + try: + s = bytes(v).decode("utf-8", errors="replace") + except Exception: + return None + s = s.strip() + return s if s else None + s = str(v).strip() + return s if s else None + + +def _sql_text_optional_jsonish(v: Any) -> str | None: + """Normalize TEXT/JSON columns across SQLite + PostgreSQL drivers (bytes/memoryview/dict).""" + if v is None: + return None + if isinstance(v, (bytes, bytearray, memoryview)): + try: + s = bytes(v).decode("utf-8", errors="replace") + except Exception: + return None + return s if s.strip() else None + if isinstance(v, (dict, list)): + return json.dumps(v, ensure_ascii=False, default=str) + s = str(v).strip() + return s if s else None + + +def _row_to_chat_message(r: Mapping[str, Any]) -> ChatMessage: + return ChatMessage( + id=int(r["id"]), + session_id=str(r["session_id"]), + role=str(r["role"]), + content=_sql_text_required(r.get("content")), + tool_calls=_sql_text_optional_jsonish(r.get("tool_calls")), + attachments=_sql_text_optional_jsonish(r.get("attachments")), + turn_uuid=_sql_text_optional_plain(r.get("turn_uuid")), + event_type=_sql_text_optional_plain(r.get("event_type")), + event_payload=_sql_text_optional_jsonish(r.get("event_payload")), + timestamp=str(r["timestamp"]), + ) + + +class ChatMessagesSaRepository: + """Phase-4 SA migration: chat_message CRUD + session last_message_at touch.""" + + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_message_and_touch_session( + self, + *, + session_id: str, + role: str, + content: str, + tool_calls: str | None, + attachments: str | None, + turn_uuid: str | None, + event_type: str | None, + event_payload: str | None, + timestamp: str, + ) -> int: + sid = str(session_id if session_id is not None else "") + with self._engine.begin() as conn: + stmt = ( + insert(chat_message) + .values( + session_id=sid, + role=str(role), + content=str(content), + tool_calls=tool_calls, + attachments=attachments, + turn_uuid=turn_uuid, + event_type=event_type, + event_payload=event_payload, + timestamp=str(timestamp), + ) + .returning(chat_message.c.id) + ) + msg_id = int(conn.execute(stmt).scalar_one()) + conn.execute( + update(chat_session) + .where(chat_session.c.id == sid) + .values(last_message_at=str(timestamp)) + ) + return msg_id + + def delete_message_and_refresh_session(self, *, session_id: str, message_id: int) -> bool: + sid = str(session_id or "").strip() + mid = int(message_id or 0) + if not sid or mid <= 0: + return False + with self._engine.begin() as conn: + res = conn.execute( + delete(chat_message).where( + chat_message.c.session_id == sid, + chat_message.c.id == mid, + ) + ) + if int(res.rowcount or 0) <= 0: + return False + last_ts = conn.execute( + select(func.max(chat_message.c.timestamp)).where(chat_message.c.session_id == sid) + ).scalar_one_or_none() + last_s = str(last_ts or "").strip() or None + conn.execute( + update(chat_session) + .where(chat_session.c.id == sid) + .values(last_message_at=last_s) + ) + return True + + def update_message_content( + self, + *, + session_id: str, + message_id: int, + content: str, + event_payload_text: str | None, + ) -> bool: + sid = str(session_id or "").strip() + mid = int(message_id or 0) + if not sid or mid <= 0: + return False + with self._engine.begin() as conn: + res = conn.execute( + update(chat_message) + .where(chat_message.c.session_id == sid, chat_message.c.id == mid) + .values( + content=str(content or ""), + event_payload=func.coalesce(literal(event_payload_text), chat_message.c.event_payload), + ) + ) + return int(res.rowcount or 0) > 0 + + def get_messages_recent_asc(self, *, session_id: str, limit: int) -> list[ChatMessage]: + if limit <= 0: + return [] + sid = str(session_id or "").strip() + if not sid: + return [] + lim = max(1, min(int(limit), 2000)) + ids_sq = ( + select(chat_message.c.id) + .where(chat_message.c.session_id == sid) + .order_by(chat_message.c.id.desc()) + .limit(lim) + .scalar_subquery() + ) + stmt = ( + select( + chat_message.c.id, + chat_message.c.session_id, + chat_message.c.role, + chat_message.c.content, + chat_message.c.tool_calls, + chat_message.c.attachments, + chat_message.c.turn_uuid, + chat_message.c.event_type, + chat_message.c.event_payload, + chat_message.c.timestamp, + ) + .where(chat_message.c.session_id == sid, chat_message.c.id.in_(ids_sq)) + .order_by(chat_message.c.id.asc()) + ) + prepended: set[int] = set() + with self._engine.connect() as conn: + rows: list[dict[str, Any]] = [dict(r) for r in conn.execute(stmt).mappings().all()] + while rows: + first = rows[0] + if str(first.get("role") or "") != "tool": + break + aid = _tool_row_assistant_message_id(first.get("tool_calls")) + if aid is None: + break + first_id = int(first["id"]) + if aid >= first_id: + break + if any(int(r["id"]) == int(aid) for r in rows): + break + if int(aid) in prepended: + break + arow = conn.execute( + select( + chat_message.c.id, + chat_message.c.session_id, + chat_message.c.role, + chat_message.c.content, + chat_message.c.tool_calls, + chat_message.c.attachments, + chat_message.c.turn_uuid, + chat_message.c.event_type, + chat_message.c.event_payload, + chat_message.c.timestamp, + ) + .where(chat_message.c.session_id == sid, chat_message.c.id == int(aid)) + .limit(1) + ).mappings().first() + if not arow: + break + prepended.add(int(aid)) + rows.insert(0, dict(arow)) + return [_row_to_chat_message(r) for r in rows] + + def get_messages_after_id( + self, *, session_id: str, after_id: int, limit: int + ) -> list[ChatMessage]: + sid = str(session_id or "").strip() + if not sid: + return [] + aid = int(after_id or 0) + lim = max(1, min(int(limit), 2000)) + stmt = ( + select( + chat_message.c.id, + chat_message.c.session_id, + chat_message.c.role, + chat_message.c.content, + chat_message.c.tool_calls, + chat_message.c.attachments, + chat_message.c.turn_uuid, + chat_message.c.event_type, + chat_message.c.event_payload, + chat_message.c.timestamp, + ) + .where(chat_message.c.session_id == sid, chat_message.c.id > aid) + .order_by(chat_message.c.id.asc()) + .limit(lim) + ) + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + return [_row_to_chat_message(dict(r)) for r in rows] + + def count_messages(self, *, session_id: str) -> int: + key = str(session_id if session_id is not None else "") + with self._engine.connect() as conn: + n = conn.execute( + select(func.count()).select_from(chat_message).where(chat_message.c.session_id == key) + ).scalar_one() + return int(n or 0) + + def session_messages_meta(self, *, session_id: str) -> SessionMessagesMeta: + key = str(session_id if session_id is not None else "") + with self._engine.connect() as conn: + row = conn.execute( + select( + func.count().label("c"), + func.max(chat_message.c.id).label("last_id"), + func.max(chat_message.c.timestamp).label("last_ts"), + ) + .where(chat_message.c.session_id == key) + ).mappings().first() + return SessionMessagesMeta( + session_id=session_id, + message_count=int(row["c"] or 0) if row else 0, + last_message_id=int(row["last_id"]) if row and row.get("last_id") is not None else None, + last_message_at=str(row["last_ts"]) if row and row.get("last_ts") is not None else None, + ) + + def last_message_id(self, *, session_id: str) -> int | None: + key = str(session_id if session_id is not None else "") + with self._engine.connect() as conn: + m = conn.execute( + select(func.max(chat_message.c.id)).where(chat_message.c.session_id == key) + ).scalar_one_or_none() + if m is None: + return None + return int(m) + + def list_messages_in_time_window( + self, *, session_id: str, start_ts: str, end_ts: str, limit: int + ) -> list[dict[str, Any]]: + sid = str(session_id or "").strip() + start = str(start_ts or "").strip() + end = str(end_ts or "").strip() + if not sid or not start or not end: + return [] + lim = max(1, min(int(limit), 2000)) + stmt = ( + select( + chat_message.c.id, + chat_message.c.session_id, + chat_message.c.role, + chat_message.c.content, + chat_message.c.tool_calls, + chat_message.c.attachments, + chat_message.c.turn_uuid, + chat_message.c.event_type, + chat_message.c.event_payload, + chat_message.c.timestamp, + ) + .where( + chat_message.c.session_id == sid, + chat_message.c.timestamp >= start, + chat_message.c.timestamp <= end, + ) + .order_by(chat_message.c.id.asc()) + .limit(lim) + ) + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + out: list[dict[str, Any]] = [] + for r in rows: + out.append( + { + "id": int(r["id"] or 0), + "session_id": str(r["session_id"] or ""), + "role": str(r["role"] or ""), + "content": str(r["content"] or ""), + "tool_calls": r["tool_calls"], + "attachments": r["attachments"], + "turn_uuid": str(r["turn_uuid"] or ""), + "event_type": str(r["event_type"] or ""), + "event_payload": r["event_payload"], + "timestamp": str(r["timestamp"] or ""), + } + ) + return out + + def delete_messages_where_session_missing(self) -> int: + """Delete ``chat_message`` rows whose ``session_id`` is not in ``chat_session`` (housekeeping).""" + sess_exists = exists( + select(1).select_from(chat_session).where(chat_session.c.id == chat_message.c.session_id) + ) + with self._engine.begin() as conn: + res = conn.execute(delete(chat_message).where(~sess_exists)) + return int(res.rowcount or 0) + + def fork_assert_anchor(self, *, source_session_id: str, up_to_message_id: int) -> None: + """Raise ``ValueError`` unless ``up_to_message_id`` exists in ``source_session_id``.""" + src = str(source_session_id if source_session_id is not None else "") + cap = int(up_to_message_id or 0) + if not src or cap <= 0: + raise ValueError("message not in session") + with self._engine.connect() as conn: + chk = conn.execute( + select(chat_message.c.id) + .where(chat_message.c.session_id == src, chat_message.c.id == cap) + .limit(1) + ).first() + if not chk: + raise ValueError("message not in session") + + def fork_copy_messages_to_session( + self, + *, + source_session_id: str, + up_to_message_id: int, + new_session_id: str, + ) -> None: + """Copy messages with ``id <= up_to_message_id`` into ``new_session_id``; remap tool assistant ids.""" + src = str(source_session_id if source_session_id is not None else "") + new_sid = str(new_session_id if new_session_id is not None else "") + cap = int(up_to_message_id or 0) + if not src or not new_sid or cap <= 0: + raise ValueError("message not in session") + with self._engine.begin() as conn: + rows_list = list( + conn.execute( + select( + chat_message.c.id, + chat_message.c.role, + chat_message.c.content, + chat_message.c.tool_calls, + chat_message.c.attachments, + chat_message.c.turn_uuid, + chat_message.c.event_type, + chat_message.c.event_payload, + chat_message.c.timestamp, + ) + .where(chat_message.c.session_id == src, chat_message.c.id <= cap) + .order_by(chat_message.c.id.asc()) + ).mappings().all() + ) + if not rows_list: + raise ValueError("message not in session") + id_map: dict[int, int] = {} + for r in rows_list: + old_id = int(r["id"]) + role = str(r["role"]) + tool_calls_text = r["tool_calls"] + if role == "tool" and tool_calls_text: + try: + meta = json.loads(str(tool_calls_text)) + if isinstance(meta, dict): + aid = meta.get("assistant_message_id") + if aid is not None: + new_aid = id_map.get(int(aid)) + if new_aid is not None: + meta = {**meta, "assistant_message_id": new_aid} + tool_calls_text = json.dumps(meta, ensure_ascii=False) + except (json.JSONDecodeError, TypeError, ValueError): + pass + stmt = ( + insert(chat_message) + .values( + session_id=new_sid, + role=role, + content=r["content"], + tool_calls=tool_calls_text, + attachments=r["attachments"], + turn_uuid=r["turn_uuid"], + event_type=r["event_type"], + event_payload=r["event_payload"], + timestamp=r["timestamp"], + ) + .returning(chat_message.c.id) + ) + new_id = int(conn.execute(stmt).scalar_one()) + id_map[old_id] = new_id + last_ts = rows_list[-1]["timestamp"] if rows_list else utc_now_iso() + conn.execute( + update(chat_session) + .where(chat_session.c.id == new_sid) + .values(last_message_at=last_ts) + ) + + def trim_messages_keep_last(self, *, session_id: str, keep_last: int) -> None: + """Delete older messages so at least ``keep_last`` newest rows remain (tool/assistant boundary aware).""" + key = str(session_id if session_id is not None else "") + kl = int(keep_last) + if not key or kl <= 0: + return + with self._engine.connect() as conn: + rows_list = [ + dict(r) + for r in conn.execute( + select(chat_message.c.id, chat_message.c.role, chat_message.c.tool_calls) + .where(chat_message.c.session_id == key) + .order_by(chat_message.c.id.asc()) + ).mappings().all() + ] + start = _trim_messages_start_index(rows_list, kl) + if start is None: + return + min_keep_id = int(rows_list[start]["id"]) + with self._engine.begin() as conn: + conn.execute( + delete(chat_message).where( + chat_message.c.session_id == key, + chat_message.c.id < min_keep_id, + ) + ) + + +__all__ = ["ChatMessagesSaRepository"] diff --git a/svc/persistence/sa_repos/chat_sessions.py b/svc/persistence/sa_repos/chat_sessions.py new file mode 100644 index 00000000..8a4dd921 --- /dev/null +++ b/svc/persistence/sa_repos/chat_sessions.py @@ -0,0 +1,424 @@ +"""chat_session (+ ui_session_owner joins) via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any, Mapping + +from sqlalchemy import and_, delete, distinct, func, insert, literal, or_, select, update +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import app_user, chat_message, chat_session, ui_session_owner +from svc.persistence.sqlite_store import ChatSession, SessionsListMeta + + +def _session_from_row(row: Mapping[str, Any]) -> ChatSession: + return ChatSession( + id=str(row["id"]), + title=str(row["title"]), + created_at=str(row["created_at"]), + last_message_at=row["last_message_at"], + ) + + +def _activity_order() -> tuple[Any, Any]: + return ( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at).desc(), + chat_session.c.created_at.desc(), + ) + + +class ChatSessionsSaRepository: + """Phase-3 SA migration: chat session rows and list queries (messages stay raw SQL for now).""" + + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_chat_session(self, *, session_id: str, title: str, created_at: str) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(chat_session).values( + id=str(session_id), + title=str(title), + created_at=str(created_at), + last_message_at=None, + ) + ) + + def fetch_chat_session_by_id(self, *, session_id: str) -> ChatSession | None: + sid = str(session_id or "").strip() + if not sid: + return None + with self._engine.connect() as conn: + row = conn.execute( + select( + chat_session.c.id, + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ) + .where(chat_session.c.id == sid) + .limit(1) + ).mappings().first() + return _session_from_row(row) if row else None + + def fetch_chat_session_for_user( + self, *, session_id: str, tenant_id: str, user_id: str + ) -> ChatSession | None: + sid = str(session_id or "").strip() + if not sid: + return None + tid, uid = str(tenant_id), str(user_id) + with self._engine.connect() as conn: + row = conn.execute( + select( + chat_session.c.id, + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ) + .select_from( + chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + ) + .where( + chat_session.c.id == sid, + ui_session_owner.c.tenant_id == tid, + ui_session_owner.c.user_id == uid, + ) + .limit(1) + ).mappings().first() + return _session_from_row(row) if row else None + + def list_chat_sessions_global(self, *, limit: int | None, offset: int) -> list[ChatSession]: + stmt = select( + chat_session.c.id, + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ).order_by(*_activity_order()) + if limit is not None: + stmt = stmt.limit(int(limit)).offset(int(offset)) + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + return [_session_from_row(r) for r in rows] + + def list_chat_sessions_for_user( + self, + *, + tenant_id: str, + user_id: str, + limit: int | None, + offset: int, + ) -> list[ChatSession]: + tid, uid = str(tenant_id), str(user_id) + stmt = ( + select( + chat_session.c.id, + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ) + .select_from( + chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + ) + .where( + ui_session_owner.c.tenant_id == tid, + ui_session_owner.c.user_id == uid, + ) + .order_by(*_activity_order()) + ) + if limit is not None: + stmt = stmt.limit(int(limit)).offset(int(offset)) + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + return [_session_from_row(r) for r in rows] + + def list_chat_sessions_for_tenant( + self, + *, + tenant_id: str, + limit: int | None, + offset: int, + ) -> list[ChatSession]: + tid = str(tenant_id) + stmt = ( + select( + chat_session.c.id, + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ) + .select_from( + chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + ) + .where(ui_session_owner.c.tenant_id == tid) + .distinct() + .order_by(*_activity_order()) + ) + if limit is not None: + stmt = stmt.limit(int(limit)).offset(int(offset)) + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + return [_session_from_row(r) for r in rows] + + def count_chat_sessions_global(self) -> int: + with self._engine.connect() as conn: + n = conn.execute(select(func.count()).select_from(chat_session)).scalar_one() + return int(n or 0) + + def sessions_list_meta_global(self) -> SessionsListMeta: + with self._engine.connect() as conn: + row = conn.execute( + select( + func.count().label("c"), + func.max( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at) + ).label("latest_activity_at"), + ).select_from(chat_session) + ).mappings().first() + return SessionsListMeta( + session_count=int(row["c"] or 0) if row else 0, + latest_activity_at=str(row["latest_activity_at"]) + if row and row.get("latest_activity_at") is not None + else None, + ) + + def sessions_list_meta_for_user(self, *, tenant_id: str, user_id: str) -> SessionsListMeta: + tid, uid = str(tenant_id), str(user_id) + with self._engine.connect() as conn: + row = conn.execute( + select( + func.count().label("c"), + func.max( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at) + ).label("latest_activity_at"), + ) + .select_from( + chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + ) + .where( + ui_session_owner.c.tenant_id == tid, + ui_session_owner.c.user_id == uid, + ) + ).mappings().first() + return SessionsListMeta( + session_count=int(row["c"] or 0) if row else 0, + latest_activity_at=str(row["latest_activity_at"]) + if row and row.get("latest_activity_at") is not None + else None, + ) + + def sessions_list_meta_for_tenant(self, *, tenant_id: str) -> SessionsListMeta: + tid = str(tenant_id) + with self._engine.connect() as conn: + row = conn.execute( + select( + func.count(distinct(chat_session.c.id)).label("c"), + func.max( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at) + ).label("latest_activity_at"), + ) + .select_from( + chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + ) + .where(ui_session_owner.c.tenant_id == tid) + ).mappings().first() + return SessionsListMeta( + session_count=int(row["c"] or 0) if row else 0, + latest_activity_at=str(row["latest_activity_at"]) + if row and row.get("latest_activity_at") is not None + else None, + ) + + def fetch_chat_session_in_tenant(self, *, session_id: str, tenant_id: str) -> ChatSession | None: + sid, tid = str(session_id or "").strip(), str(tenant_id) + if not sid: + return None + with self._engine.connect() as conn: + row = conn.execute( + select( + chat_session.c.id, + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ) + .select_from( + chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ) + ) + .where( + chat_session.c.id == sid, + ui_session_owner.c.tenant_id == tid, + ) + .limit(1) + ).mappings().first() + return _session_from_row(row) if row else None + + def list_admin_sessions( + self, + *, + tenant_id: str, + user_id: str | None, + search_lower: str | None, + active_only: bool, + active_cutoff_iso: str, + limit: int, + offset: int, + ) -> tuple[int, list[dict[str, Any]]]: + """Tenant-scoped admin session browser (parity with raw ``list_admin_sessions`` SQL).""" + tid = str(tenant_id or "").strip() + if not tid: + return 0, [] + uid = str(user_id or "").strip() or None + q_text = str(search_lower or "").strip().lower() or None + lim = max(1, min(int(limit), 500)) + off = max(0, int(offset)) + + msg_cnt = ( + select(func.count()) + .select_from(chat_message) + .where(chat_message.c.session_id == chat_session.c.id) + .scalar_subquery() + ) + joins = chat_session.join( + ui_session_owner, + ui_session_owner.c.session_id == chat_session.c.id, + ).outerjoin( + app_user, + (app_user.c.tenant_id == ui_session_owner.c.tenant_id) + & (app_user.c.id == ui_session_owner.c.user_id), + ) + conds: list[Any] = [ui_session_owner.c.tenant_id == tid] + if uid: + conds.append(ui_session_owner.c.user_id == uid) + if q_text: + like = f"%{q_text}%" + conds.append( + or_( + func.lower(func.coalesce(app_user.c.username, literal(""))).like(like), + func.lower(func.coalesce(app_user.c.display_name, literal(""))).like(like), + func.lower(func.coalesce(chat_session.c.title, literal(""))).like(like), + ) + ) + if active_only: + conds.append( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at) >= active_cutoff_iso + ) + wh = and_(*conds) + + data_stmt = ( + select( + chat_session.c.id.label("session_id"), + chat_session.c.title, + chat_session.c.created_at, + chat_session.c.last_message_at, + ui_session_owner.c.user_id, + func.coalesce(app_user.c.username, literal("")).label("username"), + func.coalesce(app_user.c.display_name, literal("")).label("display_name"), + msg_cnt.label("message_count"), + ) + .select_from(joins) + .where(wh) + .order_by( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at).desc(), + chat_session.c.created_at.desc(), + ) + .limit(lim) + .offset(off) + ) + cnt_stmt = select(func.count()).select_from(joins).where(wh) + + with self._engine.connect() as conn: + total = int(conn.execute(cnt_stmt).scalar_one() or 0) + rows = conn.execute(data_stmt).mappings().all() + out: list[dict[str, Any]] = [] + for r in rows: + out.append( + { + "session_id": str(r["session_id"] or ""), + "title": str(r["title"] or ""), + "created_at": str(r["created_at"] or ""), + "last_message_at": str(r["last_message_at"] or ""), + "user_id": str(r["user_id"] or ""), + "username": str(r["username"] or ""), + "display_name": str(r["display_name"] or ""), + "message_count": int(r["message_count"] or 0), + } + ) + return total, out + + def rename_chat_session(self, *, session_id: str, title: str) -> None: + with self._engine.begin() as conn: + conn.execute( + update(chat_session) + .where(chat_session.c.id == str(session_id)) + .values(title=str(title)) + ) + + def delete_chat_session_by_id(self, *, session_id: str) -> None: + sid = str(session_id) + with self._engine.begin() as conn: + conn.execute(delete(chat_session).where(chat_session.c.id == sid)) + + def try_delete_chat_session_for_tenant(self, *, session_id: str, tenant_id: str) -> bool: + sid, tid = str(session_id or "").strip(), str(tenant_id) + if not sid: + return False + with self._engine.begin() as conn: + chk = conn.execute( + select(1) + .select_from(ui_session_owner) + .where( + ui_session_owner.c.session_id == sid, + ui_session_owner.c.tenant_id == tid, + ) + .limit(1) + ).first() + if not chk: + return False + conn.execute(delete(chat_session).where(chat_session.c.id == sid)) + return True + + def try_delete_chat_session_for_user( + self, *, session_id: str, tenant_id: str, user_id: str + ) -> bool: + sid = str(session_id or "").strip() + if not sid: + return False + tid, uid = str(tenant_id), str(user_id) + with self._engine.begin() as conn: + chk = conn.execute( + select(1) + .select_from(ui_session_owner) + .where( + ui_session_owner.c.session_id == sid, + ui_session_owner.c.tenant_id == tid, + ui_session_owner.c.user_id == uid, + ) + .limit(1) + ).first() + if not chk: + return False + conn.execute(delete(chat_session).where(chat_session.c.id == sid)) + return True + + +__all__ = ["ChatSessionsSaRepository"] diff --git a/svc/persistence/sa_repos/session_tool_health.py b/svc/persistence/sa_repos/session_tool_health.py new file mode 100644 index 00000000..b5044455 --- /dev/null +++ b/svc/persistence/sa_repos/session_tool_health.py @@ -0,0 +1,77 @@ +"""Session tool health listing (admin) via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy import case, func, literal, select +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import chat_message, chat_session, tool_log + + +class SessionToolHealthSaRepository: + """``list_session_tool_health`` aggregates.""" + + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def list_session_tool_health( + self, *, session_id: str | None, limit: int + ) -> list[dict[str, Any]]: + lim = max(1, int(limit)) + msg_sq = ( + select( + chat_message.c.session_id, + func.sum(case((chat_message.c.role == literal("user"), 1), else_=0)).label("user_count"), + func.sum(case((chat_message.c.role == literal("assistant"), 1), else_=0)).label( + "assistant_count" + ), + ) + .group_by(chat_message.c.session_id) + .subquery() + ) + tl_sq = ( + select( + tool_log.c.session_id, + func.count(1).label("tool_count"), + func.sum(case((tool_log.c.tool_name.like("mcp__%"), 1), else_=0)).label( + "mcp_tool_count" + ), + func.max(tool_log.c.timestamp).label("last_tool_at"), + ) + .group_by(tool_log.c.session_id) + .subquery() + ) + stmt = ( + select( + chat_session.c.id.label("session_id"), + chat_session.c.title, + chat_session.c.last_message_at, + func.coalesce(msg_sq.c.user_count, literal(0)).label("user_count"), + func.coalesce(msg_sq.c.assistant_count, literal(0)).label("assistant_count"), + func.coalesce(tl_sq.c.tool_count, literal(0)).label("tool_count"), + func.coalesce(tl_sq.c.mcp_tool_count, literal(0)).label("mcp_tool_count"), + func.coalesce(tl_sq.c.last_tool_at, literal("")).label("last_tool_at"), + ) + .select_from( + chat_session.outerjoin(msg_sq, msg_sq.c.session_id == chat_session.c.id).outerjoin( + tl_sq, tl_sq.c.session_id == chat_session.c.id + ) + ) + .order_by( + func.coalesce(chat_session.c.last_message_at, chat_session.c.created_at).desc(), + ) + .limit(lim) + ) + sid = str(session_id or "").strip() + if sid: + stmt = stmt.where(chat_session.c.id == sid) + with self._engine.connect() as conn: + rows = conn.execute(stmt).mappings().all() + return [dict(r) for r in rows] + + +__all__ = ["SessionToolHealthSaRepository"] diff --git a/svc/persistence/sa_repos/tenant_bind_code.py b/svc/persistence/sa_repos/tenant_bind_code.py new file mode 100644 index 00000000..880d055a --- /dev/null +++ b/svc/persistence/sa_repos/tenant_bind_code.py @@ -0,0 +1,116 @@ +"""tenant + bind_code via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any, Mapping + +from sqlalchemy import delete, insert, select, update +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import bind_code, tenant + + +class TenantSaRepository: + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_tenant(self, *, tenant_id: str, name: str, created_at: str) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(tenant).values( + id=str(tenant_id), + name=str(name), + created_at=str(created_at), + ) + ) + + def delete_tenant(self, *, tenant_id: str) -> int: + tid = str(tenant_id or "").strip() + if not tid: + return 0 + with self._engine.begin() as conn: + res = conn.execute(delete(tenant).where(tenant.c.id == tid)) + return int(res.rowcount or 0) + + def list_tenants(self, *, limit: int) -> list[dict[str, Any]]: + lim = max(1, int(limit)) + stmt = ( + select(tenant.c.id, tenant.c.name, tenant.c.created_at) + .order_by(tenant.c.created_at.desc()) + .limit(lim) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + +class BindCodeSaRepository: + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_bind_code( + self, *, code: str, tenant_id: str, role: str, created_at: str + ) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(bind_code).values( + code=str(code), + tenant_id=str(tenant_id), + role=str(role), + created_at=str(created_at), + used_at=None, + used_by_external_user_id=None, + ) + ) + + def fetch_by_code(self, *, code: str) -> Mapping[str, Any] | None: + c = str(code or "").strip() + if not c: + return None + stmt = ( + select( + bind_code.c.code, + bind_code.c.tenant_id, + bind_code.c.role, + bind_code.c.created_at, + bind_code.c.used_at, + bind_code.c.used_by_external_user_id, + ) + .where(bind_code.c.code == c) + .limit(1) + ) + with self._engine.connect() as conn: + return conn.execute(stmt).mappings().first() + + def mark_used( + self, *, code: str, used_at: str, used_by_external_user_id: str + ) -> None: + c = str(code or "").strip() + with self._engine.begin() as conn: + conn.execute( + update(bind_code) + .where(bind_code.c.code == c) + .values(used_at=str(used_at), used_by_external_user_id=str(used_by_external_user_id)) + ) + + def list_bind_codes(self, *, tenant_id: str | None, limit: int) -> list[dict[str, Any]]: + lim = max(1, int(limit)) + stmt = select( + bind_code.c.code, + bind_code.c.tenant_id, + bind_code.c.role, + bind_code.c.created_at, + bind_code.c.used_at, + bind_code.c.used_by_external_user_id, + ).order_by(bind_code.c.created_at.desc()) + if tenant_id: + stmt = stmt.where(bind_code.c.tenant_id == str(tenant_id)) + stmt = stmt.limit(lim) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + +__all__ = ["BindCodeSaRepository", "TenantSaRepository"] diff --git a/svc/persistence/sa_repos/tool_log_queries.py b/svc/persistence/sa_repos/tool_log_queries.py new file mode 100644 index 00000000..62abb182 --- /dev/null +++ b/svc/persistence/sa_repos/tool_log_queries.py @@ -0,0 +1,133 @@ +"""tool_log read paths (MCP summaries, call logs) via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy import and_, delete, exists, func, insert, select, update +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import chat_session, tool_log + + +class ToolLogQueriesSaRepository: + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_tool_log( + self, + *, + session_id: str, + tool_name: str, + specialist: str, + args: str, + result: str, + timestamp: str, + duration_ms: int | None, + ) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(tool_log).values( + session_id=str(session_id), + tool_name=str(tool_name), + specialist=str(specialist or ""), + args=str(args), + result=str(result), + timestamp=str(timestamp), + duration_ms=duration_ms, + ) + ) + + def list_tool_logs_asc(self, *, session_id: str, limit: int) -> list[dict[str, Any]]: + sid = str(session_id or "").strip() + lim = max(1, int(limit)) + stmt = ( + select( + tool_log.c.tool_name, + tool_log.c.specialist, + tool_log.c.args, + tool_log.c.result, + tool_log.c.timestamp, + tool_log.c.duration_ms, + ) + .where(tool_log.c.session_id == sid) + .order_by(tool_log.c.id.asc()) + .limit(lim) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + def move_tool_logs_between_sessions(self, *, from_session_id: str, to_session_id: str) -> int: + src = str(from_session_id or "").strip() + dst = str(to_session_id or "").strip() + if not src or not dst or src == dst: + return 0 + with self._engine.begin() as conn: + res = conn.execute( + update(tool_log) + .where(tool_log.c.session_id == src) + .values(session_id=dst) + ) + return int(res.rowcount or 0) + + def delete_tool_logs_where_session_missing(self) -> int: + """Delete ``tool_log`` rows whose ``session_id`` is not in ``chat_session`` (housekeeping).""" + sess_exists = exists( + select(1).select_from(chat_session).where(chat_session.c.id == tool_log.c.session_id) + ) + with self._engine.begin() as conn: + res = conn.execute(delete(tool_log).where(~sess_exists)) + return int(res.rowcount or 0) + + def list_mcp_tool_usage_summary(self, *, limit: int) -> list[dict[str, Any]]: + lim = max(1, int(limit)) + n = func.count(1).label("n") + last_ts = func.max(tool_log.c.timestamp).label("last_ts") + stmt = ( + select(tool_log.c.tool_name, tool_log.c.specialist, n, last_ts) + .where(tool_log.c.tool_name.like("mcp__%")) + .group_by(tool_log.c.tool_name, tool_log.c.specialist) + .order_by(n.desc(), last_ts.desc()) + .limit(lim) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + def list_mcp_tool_aggregate_usage(self) -> list[dict[str, Any]]: + n = func.count(1).label("n") + last_ts = func.max(tool_log.c.timestamp).label("last_ts") + stmt = ( + select(tool_log.c.tool_name, n, last_ts) + .where(tool_log.c.tool_name.like("mcp__%")) + .group_by(tool_log.c.tool_name) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + def list_mcp_tool_call_logs(self, *, server_id: str | None, limit: int) -> list[dict[str, Any]]: + lim = max(1, int(limit)) + sid = str(server_id or "").strip() + conds = [tool_log.c.tool_name.like("mcp__%")] + if sid: + conds.append(tool_log.c.tool_name.like(f"mcp__{sid}__%")) + stmt = ( + select( + tool_log.c.session_id, + tool_log.c.tool_name, + tool_log.c.specialist, + tool_log.c.args, + tool_log.c.result, + tool_log.c.timestamp, + tool_log.c.duration_ms, + ) + .where(and_(*conds)) + .order_by(tool_log.c.id.desc()) + .limit(lim) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + +__all__ = ["ToolLogQueriesSaRepository"] diff --git a/svc/persistence/sa_repos/trace_events.py b/svc/persistence/sa_repos/trace_events.py new file mode 100644 index 00000000..a3c90d2a --- /dev/null +++ b/svc/persistence/sa_repos/trace_events.py @@ -0,0 +1,108 @@ +"""trace_event insert + list queries via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy import insert, select +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import trace_event + + +class TraceEventsSaRepository: + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def insert_one( + self, + *, + session_id: str, + trace_id: str, + span_id: str, + parent_span_id: str | None, + event_type: str, + payload: str, + timestamp: str, + ) -> None: + with self._engine.begin() as conn: + conn.execute( + insert(trace_event).values( + session_id=str(session_id), + trace_id=str(trace_id), + span_id=str(span_id), + parent_span_id=parent_span_id, + event_type=str(event_type), + payload=str(payload), + timestamp=str(timestamp), + ) + ) + + def insert_many(self, rows: list[dict[str, Any]]) -> None: + if not rows: + return + with self._engine.begin() as conn: + conn.execute(insert(trace_event), rows) + + def list_trace_events_desc(self, *, session_id: str, limit: int) -> list[dict[str, Any]]: + sid = str(session_id or "").strip() + lim = max(1, int(limit)) + stmt = ( + select( + trace_event.c.trace_id, + trace_event.c.span_id, + trace_event.c.parent_span_id, + trace_event.c.event_type, + trace_event.c.payload, + trace_event.c.timestamp, + ) + .where(trace_event.c.session_id == sid) + .order_by(trace_event.c.id.desc()) + .limit(lim) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + def list_trace_events_for_trace_asc( + self, *, session_id: str, trace_id: str, limit: int + ) -> list[dict[str, Any]]: + sid = str(session_id or "").strip() + tid = str(trace_id or "").strip() + lim = max(1, int(limit)) + if not sid or not tid: + return [] + stmt = ( + select( + trace_event.c.trace_id, + trace_event.c.span_id, + trace_event.c.parent_span_id, + trace_event.c.event_type, + trace_event.c.payload, + trace_event.c.timestamp, + ) + .where(trace_event.c.session_id == sid, trace_event.c.trace_id == tid) + .order_by(trace_event.c.id.asc()) + .limit(lim) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + def list_event_type_timestamp_for_trace( + self, *, session_id: str, trace_id: str + ) -> list[dict[str, Any]]: + sid = str(session_id or "").strip() + tid = str(trace_id or "").strip() + if not sid or not tid: + return [] + stmt = ( + select(trace_event.c.event_type, trace_event.c.timestamp) + .where(trace_event.c.session_id == sid, trace_event.c.trace_id == tid) + .order_by(trace_event.c.id.asc()) + ) + with self._engine.connect() as conn: + return [dict(r) for r in conn.execute(stmt).mappings().all()] + + +__all__ = ["TraceEventsSaRepository"] diff --git a/svc/persistence/sa_repos/ui_session_owner.py b/svc/persistence/sa_repos/ui_session_owner.py new file mode 100644 index 00000000..278545b7 --- /dev/null +++ b/svc/persistence/sa_repos/ui_session_owner.py @@ -0,0 +1,193 @@ +"""ui_session_owner upserts + backfills via SQLAlchemy Core.""" + +from __future__ import annotations + +from typing import Any, Mapping + +from sqlalchemy import and_, exists, func, insert, literal, select +from sqlalchemy.engine import Engine + +from svc.persistence.db.tables import ( + channel_identity_v2, + channel_session_v2, + chat_session, + ui_session_owner, +) + + +class UiSessionOwnerSaRepository: + __slots__ = ("_engine",) + + def __init__(self, engine: Engine) -> None: + self._engine = engine + + def upsert_replace( + self, *, session_id: str, tenant_id: str, user_id: str, created_at: str + ) -> None: + sid = str(session_id or "").strip() + vals = { + "session_id": sid, + "tenant_id": str(tenant_id), + "user_id": str(user_id), + "created_at": str(created_at), + } + dialect = self._engine.dialect.name + if dialect == "sqlite": + from sqlalchemy.dialects.sqlite import insert as dialect_insert + else: + from sqlalchemy.dialects.postgresql import insert as dialect_insert + + ins = dialect_insert(ui_session_owner).values(**vals) + stmt = ins.on_conflict_do_update( + index_elements=[ui_session_owner.c.session_id], + set_={ + "tenant_id": ins.excluded.tenant_id, + "user_id": ins.excluded.user_id, + "created_at": ins.excluded.created_at, + }, + ) + with self._engine.begin() as conn: + conn.execute(stmt) + + def insert_ignore( + self, *, session_id: str, tenant_id: str, user_id: str, created_at: str + ) -> None: + sid = str(session_id or "").strip() + vals = { + "session_id": sid, + "tenant_id": str(tenant_id), + "user_id": str(user_id), + "created_at": str(created_at), + } + dialect = self._engine.dialect.name + if dialect == "sqlite": + from sqlalchemy.dialects.sqlite import insert as dialect_insert + else: + from sqlalchemy.dialects.postgresql import insert as dialect_insert + + ins = dialect_insert(ui_session_owner).values(**vals) + stmt = ins.on_conflict_do_nothing(index_elements=[ui_session_owner.c.session_id]) + with self._engine.begin() as conn: + conn.execute(stmt) + + def fetch_by_session_id(self, *, session_id: str) -> Mapping[str, Any] | None: + sid = str(session_id or "").strip() + if not sid: + return None + stmt = ( + select( + ui_session_owner.c.tenant_id, + ui_session_owner.c.user_id, + ui_session_owner.c.created_at, + ) + .where(ui_session_owner.c.session_id == sid) + .limit(1) + ) + with self._engine.connect() as conn: + row = conn.execute(stmt).mappings().first() + return row + + def backfill_orphan_sessions_for_user( + self, *, tenant_id: str, user_id: str, default_created_at: str + ) -> int: + tid = str(tenant_id) + uid = str(user_id) + ts = str(default_created_at) + owned = exists( + select(1).select_from(ui_session_owner).where(ui_session_owner.c.session_id == chat_session.c.id) + ) + sel = ( + select( + chat_session.c.id.label("session_id"), + literal(tid).label("tenant_id"), + literal(uid).label("user_id"), + func.coalesce(chat_session.c.created_at, literal(ts)).label("created_at"), + ) + .where(~owned) + ) + dialect = self._engine.dialect.name + with self._engine.begin() as conn: + if dialect == "sqlite": + stmt = insert(ui_session_owner).prefix_with("OR IGNORE").from_select( + [ + ui_session_owner.c.session_id, + ui_session_owner.c.tenant_id, + ui_session_owner.c.user_id, + ui_session_owner.c.created_at, + ], + sel, + ) + else: + from sqlalchemy.dialects.postgresql import insert as pg_insert + + stmt = ( + pg_insert(ui_session_owner) + .from_select( + [ + ui_session_owner.c.session_id, + ui_session_owner.c.tenant_id, + ui_session_owner.c.user_id, + ui_session_owner.c.created_at, + ], + sel, + ) + .on_conflict_do_nothing(index_elements=[ui_session_owner.c.session_id]) + ) + res = conn.execute(stmt) + return int(res.rowcount or 0) + + def backfill_from_channel_v2(self, *, created_at: str) -> int: + ts = str(created_at) + join_on = and_( + channel_identity_v2.c.tenant_id == channel_session_v2.c.tenant_id, + channel_identity_v2.c.channel == channel_session_v2.c.channel, + channel_identity_v2.c.account_id == channel_session_v2.c.account_id, + channel_identity_v2.c.external_user_id == channel_session_v2.c.external_user_id, + ) + sel = ( + select( + channel_session_v2.c.session_id, + channel_session_v2.c.tenant_id, + channel_identity_v2.c.user_id, + literal(ts).label("created_at"), + ) + .distinct() + .select_from(channel_session_v2.join(channel_identity_v2, join_on)) + .where( + channel_session_v2.c.session_id.isnot(None), + channel_session_v2.c.session_id != "", + ) + ) + dialect = self._engine.dialect.name + with self._engine.begin() as conn: + if dialect == "sqlite": + stmt = insert(ui_session_owner).prefix_with("OR IGNORE").from_select( + [ + ui_session_owner.c.session_id, + ui_session_owner.c.tenant_id, + ui_session_owner.c.user_id, + ui_session_owner.c.created_at, + ], + sel, + ) + else: + from sqlalchemy.dialects.postgresql import insert as pg_insert + + stmt = ( + pg_insert(ui_session_owner) + .from_select( + [ + ui_session_owner.c.session_id, + ui_session_owner.c.tenant_id, + ui_session_owner.c.user_id, + ui_session_owner.c.created_at, + ], + sel, + ) + .on_conflict_do_nothing(index_elements=[ui_session_owner.c.session_id]) + ) + res = conn.execute(stmt) + return int(res.rowcount or 0) + + +__all__ = ["UiSessionOwnerSaRepository"] diff --git a/svc/persistence/sqlite_store.py b/svc/persistence/sqlite_store.py index 069ae8e3..4c208dce 100644 --- a/svc/persistence/sqlite_store.py +++ b/svc/persistence/sqlite_store.py @@ -2,17 +2,25 @@ from __future__ import annotations import base64 import json +import logging import os import re import sqlite3 import sys import hashlib import uuid +from collections.abc import Iterator +from contextlib import contextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any, Optional +from sqlalchemy.exc import IntegrityError + +from svc.persistence.pg_adapter import PgConnShim, connect_postgres +from svc.persistence.pg_compat import scrub_nul_bytes_from_jsonable, scrub_nul_bytes_from_text + if sys.platform == "win32": import ctypes from ctypes import wintypes @@ -22,6 +30,14 @@ def utc_now_iso() -> str: return datetime.now(timezone.utc).isoformat() +_LOG = logging.getLogger(__name__) + + +def _chat_message_persist_log_enabled() -> bool: + v = str(os.getenv("AIA_LOG_CHAT_MESSAGE_PERSIST") or "").strip().lower() + return v in ("1", "true", "yes", "on") + + # 内置模型配置(不可删除;rule 不在设置界面展示) LLM_BUILTIN_OLLAMA_PROFILE_ID = "00000001-0000-4000-8000-000000000001" LLM_BUILTIN_RULE_PROFILE_ID = "00000001-0000-4000-8000-000000000002" @@ -59,7 +75,7 @@ def _tool_row_assistant_message_id(tool_calls_text: str | None) -> int | None: return None -def _trim_messages_start_index(rows: list[sqlite3.Row], keep_last: int) -> int | None: +def _trim_messages_start_index(rows: list[Any], keep_last: int) -> int | None: """计算应保留的起始索引,避免 tool 消息与对应 assistant 消息断链。""" n = len(rows) if n <= keep_last: @@ -298,24 +314,169 @@ class SqliteStore: slim[k] = obj.get(k) slim["preview"] = blob[: min(cap, 4000)] return slim - def __init__(self, db_path: str | Path): - self.db_path = str(db_path) + def __init__(self, db_path: str | Path | None = None, *, postgres_url: str | None = None) -> None: + if postgres_url: + self._postgres_url = str(postgres_url).strip() + self._use_pg = True + self.db_path = self._postgres_url + else: + if db_path is None: + raise TypeError("SqliteStore requires db_path unless postgres_url= is set") + self.db_path = str(db_path) + self._postgres_url = "" + self._use_pg = False self._init_db() - def _connect(self) -> sqlite3.Connection: - # timeout:数据库被锁时的等待秒数(适配 Streamlit 多会话并发)。 + @contextmanager + def _connect(self) -> Iterator[Any]: + if self._use_pg: + raw = connect_postgres(self._postgres_url) + shim = PgConnShim(raw) + try: + yield shim + raw.commit() + except BaseException: + raw.rollback() + raise + finally: + raw.close() + return conn = sqlite3.connect(self.db_path, timeout=30.0) conn.row_factory = sqlite3.Row conn.execute("PRAGMA foreign_keys = ON;") - # WAL:读不会阻塞写,提升并发访问体验(约 10 用户规模)。 conn.execute("PRAGMA journal_mode = WAL;") - # 在 WAL 下,NORMAL 在安全性与写入吞吐之间更平衡(checkpoint 时 fsync)。 conn.execute("PRAGMA synchronous = NORMAL;") - # SQLITE_BUSY 的重试毫秒数(在连接超时之外额外生效)。 conn.execute("PRAGMA busy_timeout = 30000;") - return conn + try: + with conn: + yield conn + finally: + conn.close() + + def _shared_sa_engine(self): + """SQLAlchemy engine for the same database as :meth:`_connect` (per SQLite file or PostgreSQL URL).""" + from svc.persistence.db.engine import engine_for_sqlite_file, get_assistant_engine + + if self._use_pg: + return get_assistant_engine() + return engine_for_sqlite_file(str(self.db_path)) + + def _app_settings_repo(self): + """Lazily construct SA repository for ``app_setting`` (same DB URL as raw :meth:`_connect`).""" + from svc.persistence.sa_repos.app_settings import AppSettingsSaRepository + + r = self.__dict__.get("_app_settings_sa") + if r is None: + r = AppSettingsSaRepository(self._shared_sa_engine()) + self.__dict__["_app_settings_sa"] = r + return r + + def _auth_sessions_repo(self): + """Lazily construct SA repository for ``auth_session`` (same DB URL as raw :meth:`_connect`).""" + from svc.persistence.sa_repos.auth_sessions import AuthSessionsSaRepository + + r = self.__dict__.get("_auth_sessions_sa") + if r is None: + r = AuthSessionsSaRepository(self._shared_sa_engine()) + self.__dict__["_auth_sessions_sa"] = r + return r + + def _chat_sessions_repo(self): + """Lazily construct SA repository for ``chat_session`` / ``ui_session_owner`` list paths.""" + from svc.persistence.sa_repos.chat_sessions import ChatSessionsSaRepository + + r = self.__dict__.get("_chat_sessions_sa") + if r is None: + r = ChatSessionsSaRepository(self._shared_sa_engine()) + self.__dict__["_chat_sessions_sa"] = r + return r + + def _ui_session_owner_repo(self): + from svc.persistence.sa_repos.ui_session_owner import UiSessionOwnerSaRepository + + r = self.__dict__.get("_ui_session_owner_sa") + if r is None: + r = UiSessionOwnerSaRepository(self._shared_sa_engine()) + self.__dict__["_ui_session_owner_sa"] = r + return r + + def _tenant_repo(self): + from svc.persistence.sa_repos.tenant_bind_code import TenantSaRepository + + r = self.__dict__.get("_tenant_sa") + if r is None: + r = TenantSaRepository(self._shared_sa_engine()) + self.__dict__["_tenant_sa"] = r + return r + + def _bind_code_repo(self): + from svc.persistence.sa_repos.tenant_bind_code import BindCodeSaRepository + + r = self.__dict__.get("_bind_code_sa") + if r is None: + r = BindCodeSaRepository(self._shared_sa_engine()) + self.__dict__["_bind_code_sa"] = r + return r + + def _app_users_repo(self): + from svc.persistence.sa_repos.app_users import AppUsersSaRepository + + r = self.__dict__.get("_app_users_sa") + if r is None: + r = AppUsersSaRepository(self._shared_sa_engine()) + self.__dict__["_app_users_sa"] = r + return r + + def _chat_messages_repo(self): + """Lazily construct SA repository for ``chat_message`` hot paths.""" + from svc.persistence.sa_repos.chat_messages import ChatMessagesSaRepository + + r = self.__dict__.get("_chat_messages_sa") + if r is None: + r = ChatMessagesSaRepository(self._shared_sa_engine()) + self.__dict__["_chat_messages_sa"] = r + return r + + def _admin_user_stats_repo(self): + from svc.persistence.sa_repos.admin_user_stats import AdminUserStatsSaRepository + + r = self.__dict__.get("_admin_user_stats_sa") + if r is None: + r = AdminUserStatsSaRepository(self._shared_sa_engine()) + self.__dict__["_admin_user_stats_sa"] = r + return r + + def _session_tool_health_repo(self): + from svc.persistence.sa_repos.session_tool_health import SessionToolHealthSaRepository + + r = self.__dict__.get("_session_tool_health_sa") + if r is None: + r = SessionToolHealthSaRepository(self._shared_sa_engine()) + self.__dict__["_session_tool_health_sa"] = r + return r + + def _tool_log_queries_repo(self): + from svc.persistence.sa_repos.tool_log_queries import ToolLogQueriesSaRepository + + r = self.__dict__.get("_tool_log_queries_sa") + if r is None: + r = ToolLogQueriesSaRepository(self._shared_sa_engine()) + self.__dict__["_tool_log_queries_sa"] = r + return r + + def _trace_events_repo(self): + from svc.persistence.sa_repos.trace_events import TraceEventsSaRepository + + r = self.__dict__.get("_trace_events_sa") + if r is None: + r = TraceEventsSaRepository(self._shared_sa_engine()) + self.__dict__["_trace_events_sa"] = r + return r def _init_db(self) -> None: + if self._use_pg: + self._init_db_postgresql() + return Path(self.db_path).parent.mkdir(parents=True, exist_ok=True) with self._connect() as conn: conn.execute( @@ -1001,10 +1162,39 @@ class SqliteStore: # 会留下指向已不存在会话的行。每次初始化库时做一次幂等清理。 self._prune_rows_for_missing_chat_session(conn) - def _prune_rows_for_missing_chat_session(self, conn: sqlite3.Connection) -> None: + self._prune_orphan_chat_message_and_tool_log_outside_db_transaction() + + def _init_db_postgresql(self) -> None: + """PostgreSQL: tables from Alembic migration; run seeds and housekeeping.""" + with self._connect() as conn: + self._seed_builtin_llm_profiles(conn) + self._seed_default_permissions(conn) + conn.execute( + """ + UPDATE llm_profile SET model = ?, updated_at = ? + WHERE id = ? AND (model IS NULL OR model = 'llama3.2') + """, + (_BUILTIN_OLLAMA_MODEL_SEED, utc_now_iso(), LLM_BUILTIN_OLLAMA_PROFILE_ID), + ) + self._prune_rows_for_missing_chat_session(conn) + + # PostgreSQL: ``chat_message`` / ``tool_log`` reference ``chat_session`` with FK. + # Do **not** run :meth:`_prune_orphan_chat_message_and_tool_log_outside_db_transaction` on + # every process start: the legacy ``NOT IN (SELECT id FROM chat_session)`` housekeeping + # can race concurrent inserts or interact badly with NULL / visibility, making admin chat + # ``loadMessages`` return empty right after a turn (bubbles flash then vanish). + + def _prune_orphan_chat_message_and_tool_log_outside_db_transaction(self) -> None: + """Remove orphan ``chat_message`` / ``tool_log`` rows (must not run inside raw ``_connect`` on SQLite). + + SQLAlchemy uses its own pooled connection; running these deletes while a raw sqlite3 write + transaction is open would deadlock with ``database is locked``. + """ + self._chat_messages_repo().delete_messages_where_session_missing() + self._tool_log_queries_repo().delete_tool_logs_where_session_missing() + + def _prune_rows_for_missing_chat_session(self, conn: Any) -> None: sid_alive = "(SELECT id FROM chat_session)" - conn.execute(f"DELETE FROM chat_message WHERE session_id NOT IN {sid_alive}") - conn.execute(f"DELETE FROM tool_log WHERE session_id NOT IN {sid_alive}") conn.execute( "DELETE FROM oclaw_attempt WHERE run_id NOT IN (SELECT run_id FROM oclaw_run) " f"OR run_id IN (SELECT run_id FROM oclaw_run WHERE session_id NOT IN {sid_alive})" @@ -1013,7 +1203,7 @@ class SqliteStore: conn.execute(f"DELETE FROM oclaw_task WHERE session_id NOT IN {sid_alive}") conn.execute(f"DELETE FROM attachment_acl WHERE session_id NOT IN {sid_alive}") - def _seed_builtin_llm_profiles(self, conn: sqlite3.Connection) -> None: + def _seed_builtin_llm_profiles(self, conn: Any) -> None: ts = utc_now_iso() conn.execute( """ @@ -1030,7 +1220,7 @@ class SqliteStore: ), ) - def _seed_default_permissions(self, conn: sqlite3.Connection) -> None: + def _seed_default_permissions(self, conn: Any) -> None: ts = utc_now_iso() defaults = { "owner": { @@ -1086,62 +1276,55 @@ class SqliteStore: def create_session(self, title: str) -> ChatSession: session_id = uuid.uuid4().hex created_at = utc_now_iso() - with self._connect() as conn: - conn.execute( - "INSERT INTO chat_session (id, title, created_at, last_message_at) VALUES (?, ?, ?, NULL)", - (session_id, title, created_at), - ) + self._chat_sessions_repo().insert_chat_session( + session_id=session_id, title=title, created_at=created_at + ) return ChatSession(id=session_id, title=title, created_at=created_at, last_message_at=None) def create_session_for_user(self, *, title: str, tenant_id: str, user_id: str) -> ChatSession: s = self.create_session(title) - with self._connect() as conn: - conn.execute( - """ - INSERT OR REPLACE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) - VALUES (?, ?, ?, ?) - """, - (str(s.id), str(tenant_id), str(user_id), utc_now_iso()), + try: + self._ui_session_owner_repo().upsert_replace( + session_id=str(s.id), + tenant_id=str(tenant_id), + user_id=str(user_id), + created_at=utc_now_iso(), ) + except IntegrityError: + # ``chat_session`` is already committed; without ``ui_session_owner`` the session is invisible + # to ``list_sessions_for_user`` / ``get_session_for_user`` (INNER JOIN). PostgreSQL enforces + # FK from ``ui_session_owner`` to ``tenant`` / ``app_user``; a failed owner row leaves a + # "ghost" session that looks like PG-specific data loss vs SQLite (weaker FK history). + try: + self.delete_session(str(s.id)) + except Exception: + pass + raise return s def ensure_ui_session_owner(self, *, session_id: str, tenant_id: str, user_id: str) -> None: - with self._connect() as conn: - conn.execute( - """ - INSERT OR IGNORE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) - VALUES (?, ?, ?, ?) - """, - (str(session_id), str(tenant_id), str(user_id), utc_now_iso()), - ) + self._ui_session_owner_repo().insert_ignore( + session_id=str(session_id), + tenant_id=str(tenant_id), + user_id=str(user_id), + created_at=utc_now_iso(), + ) def get_session(self, session_id: str) -> Optional[ChatSession]: - with self._connect() as conn: - row = conn.execute( - "SELECT id, title, created_at, last_message_at FROM chat_session WHERE id = ?", - (session_id,), - ).fetchone() - if not row: - return None - return ChatSession( - id=row["id"], - title=row["title"], - created_at=row["created_at"], - last_message_at=row["last_message_at"], - ) + return self._chat_sessions_repo().fetch_chat_session_by_id(session_id=session_id) def get_ui_session_owner(self, *, session_id: str) -> dict[str, Any] | None: sid = str(session_id or "").strip() if not sid: return None - with self._connect() as conn: - row = conn.execute( - "SELECT tenant_id, user_id, created_at FROM ui_session_owner WHERE session_id = ? LIMIT 1", - (sid,), - ).fetchone() + row = self._ui_session_owner_repo().fetch_by_session_id(session_id=sid) if not row: return None - return {"tenant_id": str(row["tenant_id"] or ""), "user_id": str(row["user_id"] or ""), "created_at": str(row["created_at"] or "")} + return { + "tenant_id": str(row["tenant_id"] or ""), + "user_id": str(row["user_id"] or ""), + "created_at": str(row["created_at"] or ""), + } def backfill_orphan_chat_sessions_for_user(self, *, tenant_id: str, user_id: str) -> int: """将**当前库中所有**尚无 ``ui_session_owner`` 的 ``chat_session`` 归属到指定用户。 @@ -1150,62 +1333,21 @@ class SqliteStore: 已从 HTTP 列表接口移除自动调用;仅保留供单租户数据修复时在 Python 控制台等场景**显式**调用。 """ ts = utc_now_iso() - with self._connect() as conn: - cur = conn.execute( - """ - INSERT OR IGNORE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) - SELECT s.id, ?, ?, COALESCE(s.created_at, ?) - FROM chat_session s - WHERE NOT EXISTS (SELECT 1 FROM ui_session_owner o WHERE o.session_id = s.id) - """, - (str(tenant_id), str(user_id), ts), - ) - return int(cur.rowcount or 0) + return self._ui_session_owner_repo().backfill_orphan_sessions_for_user( + tenant_id=str(tenant_id), user_id=str(user_id), default_created_at=ts + ) def get_session_for_user(self, *, session_id: str, tenant_id: str, user_id: str) -> Optional[ChatSession]: - sid = str(session_id or "").strip() - if not sid: - return None - with self._connect() as conn: - owner = conn.execute( - """ - SELECT 1 - FROM ui_session_owner - WHERE session_id = ? AND tenant_id = ? AND user_id = ? - LIMIT 1 - """, - (sid, str(tenant_id), str(user_id)), - ).fetchone() - if not owner: - return None - return self.get_session(sid) + return self._chat_sessions_repo().fetch_chat_session_for_user( + session_id=session_id, tenant_id=tenant_id, user_id=user_id + ) def list_sessions( self, limit: int | None = None, offset: int = 0, ) -> list[ChatSession]: - lim_sql = "" - params: list[Any] = [] - if limit is not None: - lim_sql = "LIMIT ? OFFSET ?" - params.extend([int(limit), int(offset)]) - sql = f""" - SELECT id, title, created_at, last_message_at FROM chat_session - ORDER BY COALESCE(last_message_at, created_at) DESC, created_at DESC - {lim_sql} - """ - with self._connect() as conn: - rows = conn.execute(sql, params).fetchall() - return [ - ChatSession( - id=r["id"], - title=r["title"], - created_at=r["created_at"], - last_message_at=r["last_message_at"], - ) - for r in rows - ] + return self._chat_sessions_repo().list_chat_sessions_global(limit=limit, offset=int(offset)) def list_sessions_for_user( self, @@ -1215,68 +1357,22 @@ class SqliteStore: limit: int | None = None, offset: int = 0, ) -> list[ChatSession]: - lim_sql = "" - params: list[Any] = [str(tenant_id), str(user_id)] - if limit is not None: - lim_sql = "LIMIT ? OFFSET ?" - params.extend([int(limit), int(offset)]) - sql = f""" - SELECT s.id, s.title, s.created_at, s.last_message_at - FROM chat_session s - JOIN ui_session_owner o ON o.session_id = s.id - WHERE o.tenant_id = ? AND o.user_id = ? - ORDER BY COALESCE(s.last_message_at, s.created_at) DESC, s.created_at DESC - {lim_sql} - """ - with self._connect() as conn: - rows = conn.execute(sql, params).fetchall() - return [ - ChatSession( - id=r["id"], - title=r["title"], - created_at=r["created_at"], - last_message_at=r["last_message_at"], - ) - for r in rows - ] - - def count_sessions(self) -> int: - sql = "SELECT COUNT(*) AS c FROM chat_session" - with self._connect() as conn: - row = conn.execute(sql).fetchone() - return int(row["c"]) if row else 0 - - def get_sessions_list_meta(self) -> SessionsListMeta: - with self._connect() as conn: - row = conn.execute( - """ - SELECT - COUNT(*) AS c, - MAX(COALESCE(last_message_at, created_at)) AS latest_activity_at - FROM chat_session - """ - ).fetchone() - return SessionsListMeta( - session_count=int(row["c"] or 0) if row else 0, - latest_activity_at=str(row["latest_activity_at"]) if row and row["latest_activity_at"] is not None else None, + return self._chat_sessions_repo().list_chat_sessions_for_user( + tenant_id=tenant_id, + user_id=user_id, + limit=limit, + offset=int(offset), ) + def count_sessions(self) -> int: + return self._chat_sessions_repo().count_chat_sessions_global() + + def get_sessions_list_meta(self) -> SessionsListMeta: + return self._chat_sessions_repo().sessions_list_meta_global() + def get_sessions_list_meta_for_user(self, *, tenant_id: str, user_id: str) -> SessionsListMeta: - with self._connect() as conn: - row = conn.execute( - """ - SELECT - COUNT(*) AS c, - MAX(COALESCE(s.last_message_at, s.created_at)) AS latest_activity_at - FROM chat_session s - JOIN ui_session_owner o ON o.session_id = s.id - WHERE o.tenant_id = ? AND o.user_id = ? - """, - (str(tenant_id), str(user_id)), - ).fetchone() - return SessionsListMeta( - session_count=int(row["c"] or 0) if row else 0, - latest_activity_at=str(row["latest_activity_at"]) if row and row["latest_activity_at"] is not None else None, + return self._chat_sessions_repo().sessions_list_meta_for_user( + tenant_id=tenant_id, user_id=user_id ) def list_sessions_for_tenant( @@ -1287,70 +1383,17 @@ class SqliteStore: offset: int = 0, ) -> list[ChatSession]: """All chat sessions that belong to ``tenant_id`` via ``ui_session_owner`` (any user).""" - lim_sql = "" - params: list[Any] = [str(tenant_id)] - if limit is not None: - lim_sql = "LIMIT ? OFFSET ?" - params.extend([int(limit), int(offset)]) - sql = f""" - SELECT DISTINCT s.id, s.title, s.created_at, s.last_message_at - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id AND o.tenant_id = ? - ORDER BY COALESCE(s.last_message_at, s.created_at) DESC, s.created_at DESC - {lim_sql} - """ - with self._connect() as conn: - rows = conn.execute(sql, params).fetchall() - return [ - ChatSession( - id=r["id"], - title=r["title"], - created_at=r["created_at"], - last_message_at=r["last_message_at"], - ) - for r in rows - ] + return self._chat_sessions_repo().list_chat_sessions_for_tenant( + tenant_id=tenant_id, limit=limit, offset=int(offset) + ) def get_sessions_list_meta_for_tenant(self, *, tenant_id: str) -> SessionsListMeta: - with self._connect() as conn: - row = conn.execute( - """ - SELECT - COUNT(DISTINCT s.id) AS c, - MAX(COALESCE(s.last_message_at, s.created_at)) AS latest_activity_at - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id AND o.tenant_id = ? - """, - (str(tenant_id),), - ).fetchone() - return SessionsListMeta( - session_count=int(row["c"] or 0) if row else 0, - latest_activity_at=str(row["latest_activity_at"]) if row and row["latest_activity_at"] is not None else None, - ) + return self._chat_sessions_repo().sessions_list_meta_for_tenant(tenant_id=tenant_id) def get_session_in_tenant(self, *, session_id: str, tenant_id: str) -> Optional[ChatSession]: """Session exists and is linked to this tenant (``administrator`` global browse).""" - sid = str(session_id or "").strip() - if not sid: - return None - with self._connect() as conn: - row = conn.execute( - """ - SELECT s.id, s.title, s.created_at, s.last_message_at - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id AND o.tenant_id = ? - WHERE s.id = ? - LIMIT 1 - """, - (str(tenant_id), sid), - ).fetchone() - if not row: - return None - return ChatSession( - id=row["id"], - title=row["title"], - created_at=row["created_at"], - last_message_at=row["last_message_at"], + return self._chat_sessions_repo().fetch_chat_session_in_tenant( + session_id=session_id, tenant_id=tenant_id ) @staticmethod @@ -1646,71 +1689,25 @@ class SqliteStore: lim = max(1, min(int(limit), 500)) off = max(0, int(offset)) - cond = ["o.tenant_id = ?"] - params: list[Any] = [tid] - if uid: - cond.append("o.user_id = ?") - params.append(uid) - if q_text: - like = f"%{q_text}%" - cond.append( - "(LOWER(COALESCE(u.username,'')) LIKE ? OR LOWER(COALESCE(u.display_name,'')) LIKE ? " - "OR LOWER(COALESCE(s.title,'')) LIKE ?)" - ) - params.extend([like, like, like]) - if active_only: - cond.append("COALESCE(s.last_message_at, s.created_at) >= ?") - params.append(cutoff) - where_sql = " AND ".join(cond) - - with self._connect() as conn: - total_row = conn.execute( - f""" - SELECT COUNT(*) AS c - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id - LEFT JOIN app_user u ON u.tenant_id = o.tenant_id AND u.id = o.user_id - WHERE {where_sql} - """, - params, - ).fetchone() - rows = conn.execute( - f""" - SELECT - s.id AS session_id, - s.title AS title, - s.created_at AS created_at, - s.last_message_at AS last_message_at, - o.user_id AS user_id, - COALESCE(u.username, '') AS username, - COALESCE(u.display_name, '') AS display_name, - (SELECT COUNT(*) FROM chat_message m WHERE m.session_id = s.id) AS message_count - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id - LEFT JOIN app_user u ON u.tenant_id = o.tenant_id AND u.id = o.user_id - WHERE {where_sql} - ORDER BY COALESCE(s.last_message_at, s.created_at) DESC, s.created_at DESC - LIMIT ? OFFSET ? - """, - [*params, lim, off], - ).fetchall() + total, rows = self._chat_sessions_repo().list_admin_sessions( + tenant_id=tid, + user_id=uid or None, + search_lower=q_text or None, + active_only=active_only, + active_cutoff_iso=cutoff, + limit=lim, + offset=off, + ) out: list[dict[str, Any]] = [] for r in rows: last_at = str(r["last_message_at"] or r["created_at"] or "") out.append( { - "session_id": str(r["session_id"] or ""), - "title": str(r["title"] or ""), - "created_at": str(r["created_at"] or ""), - "last_message_at": str(r["last_message_at"] or ""), - "user_id": str(r["user_id"] or ""), - "username": str(r["username"] or ""), - "display_name": str(r["display_name"] or ""), - "message_count": int(r["message_count"] or 0), + **r, "is_active_30m": bool(last_at and last_at >= cutoff), } ) - return int((total_row["c"] if total_row else 0) or 0), out + return total, out def list_admin_user_stats( self, @@ -1735,55 +1732,18 @@ class SqliteStore: lim = max(1, min(int(limit), 500)) off = max(0, int(offset)) - cond = ["tenant_id = ?"] - params: list[Any] = [tid] - if q_text: - like = f"%{q_text}%" - cond.append("(LOWER(COALESCE(username,'')) LIKE ? OR LOWER(COALESCE(display_name,'')) LIKE ?)") - params.extend([like, like]) - where_sql = " AND ".join(cond) - with self._connect() as conn: - total_row = conn.execute( - f"SELECT COUNT(*) AS c FROM app_user WHERE {where_sql}", - params, - ).fetchone() - rows = conn.execute( - f""" - SELECT - id AS user_id, - username, - COALESCE(display_name, '') AS display_name, - role, - is_active - FROM app_user - WHERE {where_sql} - ORDER BY username ASC - LIMIT ? OFFSET ? - """, - [*params, lim, off], - ).fetchall() - total_active_sessions = conn.execute( - """ - SELECT COUNT(DISTINCT s.id) AS c - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id - WHERE o.tenant_id = ? AND COALESCE(s.last_message_at, s.created_at) >= ? - """, - (tid, cutoff), - ).fetchone() - total_active_logins = conn.execute( - """ - SELECT COUNT(*) AS c - FROM auth_session - WHERE tenant_id = ? - AND revoked_at IS NULL - AND expires_at > ? - AND last_seen_at >= ? - """, - (tid, cutoff, cutoff), - ).fetchone() + pack = self._admin_user_stats_repo().fetch( + tenant_id=tid, + search_lower=q_text or None, + cutoff_iso=cutoff, + limit=lim, + offset=off, + ) + total_row_c = int(pack["total_users"] or 0) + rows = pack["user_rows"] + total_active_sessions = pack["total_active_sessions_30m"] + total_active_logins = pack["total_active_logins_30m"] - user_ids = [str(r["user_id"] or "") for r in rows if str(r["user_id"] or "").strip()] token_by_user: dict[str, int] = {} sessions_count_by_user: dict[str, int] = {} active_sessions_by_user: dict[str, int] = {} @@ -1791,80 +1751,30 @@ class SqliteStore: active_logins_by_user: dict[str, int] = {} last_seen_at_by_user: dict[str, str] = {} - if user_ids: - ph = ",".join("?" for _ in user_ids) - with self._connect() as conn: - token_rows = conn.execute( - f""" - SELECT o.user_id AS user_id, e.payload AS payload - FROM trace_event e - INNER JOIN ui_session_owner o ON o.session_id = e.session_id - WHERE o.tenant_id = ? AND o.user_id IN ({ph}) - """, - [tid, *user_ids], - ).fetchall() - for tr in token_rows: - uid = str(tr["user_id"] or "") - if not uid: - continue - try: - payload = json.loads(tr["payload"] or "{}") - except Exception: - payload = {} - p = int(payload.get("prompt_tokens_est") or 0) - r2 = int(payload.get("response_tokens_est") or 0) - token_by_user[uid] = int(token_by_user.get(uid, 0)) + max(0, p) + max(0, r2) + for tr in pack["trace_rows"]: + uid = str(tr.get("user_id") or "") + if not uid: + continue + try: + payload = json.loads(tr.get("payload") or "{}") + except Exception: + payload = {} + p = int(payload.get("prompt_tokens_est") or 0) + r2 = int(payload.get("response_tokens_est") or 0) + token_by_user[uid] = int(token_by_user.get(uid, 0)) + max(0, p) + max(0, r2) - s_rows = conn.execute( - f""" - SELECT user_id, COUNT(*) AS c - FROM ui_session_owner - WHERE tenant_id = ? AND user_id IN ({ph}) - GROUP BY user_id - """, - [tid, *user_ids], - ).fetchall() - for sr in s_rows: - sessions_count_by_user[str(sr["user_id"] or "")] = int(sr["c"] or 0) + for sr in pack["sessions_count_rows"]: + sessions_count_by_user[str(sr.get("user_id") or "")] = int(sr.get("c") or 0) - sess_rows = conn.execute( - f""" - SELECT - o.user_id AS user_id, - COUNT(DISTINCT CASE WHEN COALESCE(s.last_message_at, s.created_at) >= ? THEN s.id END) AS active_30m, - MAX(COALESCE(s.last_message_at, s.created_at)) AS last_message_at - FROM chat_session s - INNER JOIN ui_session_owner o ON o.session_id = s.id - WHERE o.tenant_id = ? AND o.user_id IN ({ph}) - GROUP BY o.user_id - """, - [cutoff, tid, *user_ids], - ).fetchall() - for rr in sess_rows: - uid = str(rr["user_id"] or "") - active_sessions_by_user[uid] = int(rr["active_30m"] or 0) - last_message_at_by_user[uid] = str(rr["last_message_at"] or "") + for rr in pack["active_sess_rows"]: + uid = str(rr.get("user_id") or "") + active_sessions_by_user[uid] = int(rr.get("active_30m") or 0) + last_message_at_by_user[uid] = str(rr.get("last_message_at") or "") - login_rows = conn.execute( - f""" - SELECT - user_id AS user_id, - COUNT(*) AS c, - MAX(last_seen_at) AS last_seen_at - FROM auth_session - WHERE tenant_id = ? - AND user_id IN ({ph}) - AND revoked_at IS NULL - AND expires_at > ? - AND last_seen_at >= ? - GROUP BY user_id - """, - [tid, *user_ids, cutoff, cutoff], - ).fetchall() - for lr in login_rows: - uid = str(lr["user_id"] or "") - active_logins_by_user[uid] = int(lr["c"] or 0) - last_seen_at_by_user[uid] = str(lr["last_seen_at"] or "") + for lr in pack["login_rows"]: + uid = str(lr.get("user_id") or "") + active_logins_by_user[uid] = int(lr.get("c") or 0) + last_seen_at_by_user[uid] = str(lr.get("last_seen_at") or "") users: list[dict[str, Any]] = [] for r in rows: @@ -1887,73 +1797,30 @@ class SqliteStore: totals = { "total_tokens_est": int(sum(int(x.get("total_tokens_est") or 0) for x in users)), - "active_sessions_30m": int((total_active_sessions["c"] if total_active_sessions else 0) or 0), - "active_logins_30m": int((total_active_logins["c"] if total_active_logins else 0) or 0), - "users_count": int((total_row["c"] if total_row else 0) or 0), + "active_sessions_30m": int(total_active_sessions or 0), + "active_logins_30m": int(total_active_logins or 0), + "users_count": int(total_row_c or 0), } - return int((total_row["c"] if total_row else 0) or 0), users, totals + return int(total_row_c or 0), users, totals def delete_session_in_tenant(self, *, session_id: str, tenant_id: str) -> bool: """Delete session if it belongs to tenant (used by administrator account).""" - sid = str(session_id or "").strip() - if not sid: - return False - with self._connect() as conn: - owner = conn.execute( - """ - SELECT 1 FROM ui_session_owner - WHERE session_id = ? AND tenant_id = ? - LIMIT 1 - """, - (sid, str(tenant_id)), - ).fetchone() - if not owner: - return False - conn.execute("DELETE FROM chat_session WHERE id = ?", (sid,)) - return True + return self._chat_sessions_repo().try_delete_chat_session_for_tenant( + session_id=session_id, tenant_id=tenant_id + ) def delete_session(self, session_id: str) -> None: - with self._connect() as conn: - conn.execute("DELETE FROM chat_session WHERE id = ?", (session_id,)) + self._chat_sessions_repo().delete_chat_session_by_id(session_id=session_id) def delete_session_for_user(self, *, session_id: str, tenant_id: str, user_id: str) -> bool: - with self._connect() as conn: - owner = conn.execute( - """ - SELECT 1 - FROM ui_session_owner - WHERE session_id = ? AND tenant_id = ? AND user_id = ? - LIMIT 1 - """, - (str(session_id), str(tenant_id), str(user_id)), - ).fetchone() - if not owner: - return False - conn.execute("DELETE FROM chat_session WHERE id = ?", (str(session_id),)) - return True + return self._chat_sessions_repo().try_delete_chat_session_for_user( + session_id=session_id, tenant_id=tenant_id, user_id=user_id + ) def delete_message(self, *, session_id: str, message_id: int) -> bool: - sid = str(session_id or "").strip() - mid = int(message_id or 0) - if not sid or mid <= 0: - return False - with self._connect() as conn: - cur = conn.execute( - "DELETE FROM chat_message WHERE session_id = ? AND id = ?", - (sid, mid), - ) - if int(cur.rowcount or 0) <= 0: - return False - last_row = conn.execute( - "SELECT MAX(timestamp) AS ts FROM chat_message WHERE session_id = ?", - (sid,), - ).fetchone() - last_ts = str((last_row["ts"] if last_row else "") or "").strip() or None - conn.execute( - "UPDATE chat_session SET last_message_at = ? WHERE id = ?", - (last_ts, sid), - ) - return True + return self._chat_messages_repo().delete_message_and_refresh_session( + session_id=session_id, message_id=message_id + ) def add_message( self, @@ -1973,42 +1840,57 @@ class SqliteStore: if isinstance(tool_calls, str): tool_calls_text = tool_calls else: - tool_calls_text = json.dumps(tool_calls, ensure_ascii=False) + tool_calls_text = json.dumps( + scrub_nul_bytes_from_jsonable(tool_calls), ensure_ascii=False + ) attachments_text = None if attachments is not None: - attachments_text = json.dumps(attachments, ensure_ascii=False) + attachments_text = json.dumps( + scrub_nul_bytes_from_jsonable(attachments), ensure_ascii=False + ) event_payload_text = None if event_payload is not None: if isinstance(event_payload, str): event_payload_text = event_payload else: - event_payload_text = json.dumps(event_payload, ensure_ascii=False, default=str) + event_payload_text = json.dumps( + scrub_nul_bytes_from_jsonable(event_payload), ensure_ascii=False, default=str + ) turn_uuid_text = str(turn_uuid or "").strip() or None event_type_text = str(event_type or "").strip() or None - with self._connect() as conn: - cur = conn.execute( - """ - INSERT INTO chat_message ( - session_id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - ) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - session_id, - role, - content, - tool_calls_text, - attachments_text, - turn_uuid_text, - event_type_text, - event_payload_text, - ts, - ), - ) - msg_id = int(cur.lastrowid) - conn.execute( - "UPDATE chat_session SET last_message_at = ? WHERE id = ?", - (ts, session_id), + # PostgreSQL TEXT rejects NUL; SQLite accepts it — strip before SA insert. + content_clean = str(scrub_nul_bytes_from_text(str(content)) or "") + tool_calls_text = scrub_nul_bytes_from_text(tool_calls_text) + attachments_text = scrub_nul_bytes_from_text(attachments_text) + event_payload_text = scrub_nul_bytes_from_text(event_payload_text) + turn_uuid_text = scrub_nul_bytes_from_text(turn_uuid_text) + event_type_text = scrub_nul_bytes_from_text(event_type_text) + if turn_uuid_text is not None and not str(turn_uuid_text).strip(): + turn_uuid_text = None + if event_type_text is not None and not str(event_type_text).strip(): + event_type_text = None + msg_id = self._chat_messages_repo().insert_message_and_touch_session( + session_id=str(session_id if session_id is not None else ""), + role=str(role), + content=content_clean, + tool_calls=tool_calls_text, + attachments=attachments_text, + turn_uuid=turn_uuid_text, + event_type=event_type_text, + event_payload=event_payload_text, + timestamp=str(ts), + ) + if _chat_message_persist_log_enabled(): + _LOG.warning( + "chat_message_persisted backend=%s session_id=%s message_id=%s role=%s event_type=%s " + "content_len=%d has_tool_calls=%s", + "postgresql" if self._use_pg else "sqlite", + str(session_id or "").strip(), + int(msg_id), + str(role), + str(event_type_text or ""), + len(str(content_clean or "")), + bool(tool_calls_text), ) try: self.sync_attachment_acl_from_chat_message_attachments( @@ -2023,7 +1905,7 @@ class SqliteStore: id=msg_id, session_id=session_id, role=role, - content=content, + content=content_clean, tool_calls=tool_calls_text, attachments=attachments_text, turn_uuid=turn_uuid_text, @@ -2051,17 +1933,17 @@ class SqliteStore: if isinstance(event_payload, str): event_payload_text = event_payload else: - event_payload_text = json.dumps(event_payload, ensure_ascii=False, default=str) - with self._connect() as conn: - cur = conn.execute( - """ - UPDATE chat_message - SET content = ?, event_payload = COALESCE(?, event_payload) - WHERE session_id = ? AND id = ? - """, - (str(content or ""), event_payload_text, sid, mid), - ) - return int(cur.rowcount or 0) > 0 + event_payload_text = json.dumps( + scrub_nul_bytes_from_jsonable(event_payload), ensure_ascii=False, default=str + ) + event_payload_text = scrub_nul_bytes_from_text(event_payload_text) + content_u = str(scrub_nul_bytes_from_text(str(content or "")) or "") + return self._chat_messages_repo().update_message_content( + session_id=sid, + message_id=mid, + content=content_u, + event_payload_text=event_payload_text, + ) def get_messages(self, session_id: str, limit: int = 200) -> list[ChatMessage]: """返回最近 ``limit`` 条消息,顺序为时间正序(窗口内最早的一条在前)。""" @@ -2071,65 +1953,7 @@ class SqliteStore: if not sid: return [] lim = max(1, min(int(limit), 2000)) - with self._connect() as conn: - rows = list( - conn.execute( - """ - SELECT id, session_id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - FROM chat_message - WHERE session_id = ? AND id IN ( - SELECT id FROM chat_message WHERE session_id = ? ORDER BY id DESC LIMIT ? - ) - ORDER BY id ASC - """, - (sid, sid, lim), - ).fetchall() - ) - # Preserve tool->assistant pairing at the boundary: if the first kept row is a tool message, - # fetch and prepend the referenced assistant message (emitted tool_calls) when it's outside the window. - # This prevents OpenAI-compatible gateways from rejecting unpaired tool results. - prepended: set[int] = set() - while rows: - first = rows[0] - if str(first["role"] or "") != "tool": - break - aid = _tool_row_assistant_message_id(first["tool_calls"]) - if aid is None: - break - first_id = int(first["id"]) - if aid >= first_id: - break - if any(int(r["id"]) == int(aid) for r in rows): - break - if int(aid) in prepended: - break - arow = conn.execute( - """ - SELECT id, session_id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - FROM chat_message - WHERE session_id = ? AND id = ? - """, - (sid, int(aid)), - ).fetchone() - if not arow: - break - prepended.add(int(aid)) - rows.insert(0, arow) - return [ - ChatMessage( - id=int(r["id"]), - session_id=r["session_id"], - role=r["role"], - content=r["content"], - tool_calls=r["tool_calls"], - attachments=r["attachments"], - turn_uuid=r["turn_uuid"], - event_type=r["event_type"], - event_payload=r["event_payload"], - timestamp=r["timestamp"], - ) - for r in rows - ] + return self._chat_messages_repo().get_messages_recent_asc(session_id=sid, limit=lim) def get_messages_after_id(self, *, session_id: str, after_id: int, limit: int = 200) -> list[ChatMessage]: """Return messages with id > after_id in ASC order (bounded by limit).""" @@ -2138,32 +1962,9 @@ class SqliteStore: return [] aid = int(after_id or 0) lim = max(1, min(int(limit), 2000)) - with self._connect() as conn: - rows = conn.execute( - """ - SELECT id, session_id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - FROM chat_message - WHERE session_id = ? AND id > ? - ORDER BY id ASC - LIMIT ? - """, - (sid, aid, lim), - ).fetchall() - return [ - ChatMessage( - id=int(r["id"]), - session_id=r["session_id"], - role=r["role"], - content=r["content"], - tool_calls=r["tool_calls"], - attachments=r["attachments"], - turn_uuid=r["turn_uuid"], - event_type=r["event_type"], - event_payload=r["event_payload"], - timestamp=r["timestamp"], - ) - for r in rows - ] + return self._chat_messages_repo().get_messages_after_id( + session_id=sid, after_id=aid, limit=lim + ) def add_tool_log( self, @@ -2184,35 +1985,20 @@ class SqliteStore: cap = max(20_000, min(int(raw_cap), 2_000_000)) args_capped = self._cap_json_for_log(args, max_chars=cap, keep_keys=()) result_capped = self._cap_json_for_log(result, max_chars=cap, keep_keys=("ok", "error_code", "error")) - with self._connect() as conn: - conn.execute( - """ - INSERT INTO tool_log (session_id, tool_name, specialist, args, result, timestamp, duration_ms) - VALUES (?, ?, ?, ?, ?, ?, ?) - """, - ( - session_id, - tool_name, - str(specialist or ""), - json.dumps(args_capped, ensure_ascii=False, default=str), - json.dumps(result_capped, ensure_ascii=False, default=str), - ts, - duration_ms, - ), - ) + self._tool_log_queries_repo().insert_tool_log( + session_id=str(session_id), + tool_name=str(tool_name), + specialist=str(specialist or ""), + args=json.dumps(args_capped, ensure_ascii=False, default=str), + result=json.dumps(result_capped, ensure_ascii=False, default=str), + timestamp=str(ts), + duration_ms=duration_ms, + ) def get_tool_logs(self, session_id: str, limit: int = 200) -> list[dict[str, Any]]: - with self._connect() as conn: - rows = conn.execute( - """ - SELECT tool_name, specialist, args, result, timestamp, duration_ms - FROM tool_log - WHERE session_id = ? - ORDER BY id ASC - LIMIT ? - """, - (session_id, limit), - ).fetchall() + rows = self._tool_log_queries_repo().list_tool_logs_asc( + session_id=str(session_id), limit=max(1, int(limit)) + ) out: list[dict[str, Any]] = [] for r in rows: out.append( @@ -2228,51 +2014,11 @@ class SqliteStore: return out def list_session_tool_health(self, *, session_id: str | None = None, limit: int = 80) -> list[dict[str, Any]]: - params: list[Any] = [] - where = "" - sid = str(session_id or "").strip() - if sid: - where = "WHERE s.id = ?" - params.append(sid) - params.append(max(1, int(limit))) - with self._connect() as conn: - rows = conn.execute( - f""" - SELECT - s.id AS session_id, - s.title AS title, - s.last_message_at AS last_message_at, - COALESCE(msg.user_count, 0) AS user_count, - COALESCE(msg.assistant_count, 0) AS assistant_count, - COALESCE(tl.tool_count, 0) AS tool_count, - COALESCE(tl.mcp_tool_count, 0) AS mcp_tool_count, - COALESCE(tl.last_tool_at, '') AS last_tool_at - FROM chat_session s - LEFT JOIN ( - SELECT - session_id, - SUM(CASE WHEN role = 'user' THEN 1 ELSE 0 END) AS user_count, - SUM(CASE WHEN role = 'assistant' THEN 1 ELSE 0 END) AS assistant_count - FROM chat_message - GROUP BY session_id - ) msg ON msg.session_id = s.id - LEFT JOIN ( - SELECT - session_id, - COUNT(1) AS tool_count, - SUM(CASE WHEN tool_name LIKE 'mcp__%' THEN 1 ELSE 0 END) AS mcp_tool_count, - MAX(timestamp) AS last_tool_at - FROM tool_log - GROUP BY session_id - ) tl ON tl.session_id = s.id - {where} - ORDER BY COALESCE(s.last_message_at, s.created_at) DESC - LIMIT ? - """, - tuple(params), - ).fetchall() + lim = max(1, int(limit)) + sid = str(session_id or "").strip() or None + raw = self._session_tool_health_repo().list_session_tool_health(session_id=sid, limit=lim) out: list[dict[str, Any]] = [] - for r in rows: + for r in raw: tool_count = int(r["tool_count"] or 0) assistant_count = int(r["assistant_count"] or 0) unhealthy = assistant_count > 0 and tool_count == 0 @@ -2296,26 +2042,12 @@ class SqliteStore: dst = str(to_session_id or "").strip() if not src or not dst or src == dst: return 0 - with self._connect() as conn: - cur = conn.execute( - "UPDATE tool_log SET session_id = ? WHERE session_id = ?", - (dst, src), - ) - return int(cur.rowcount or 0) + return self._tool_log_queries_repo().move_tool_logs_between_sessions( + from_session_id=src, to_session_id=dst + ) def list_mcp_tool_usage_summary(self, *, limit: int = 200) -> list[dict[str, Any]]: - with self._connect() as conn: - rows = conn.execute( - """ - SELECT tool_name, specialist, COUNT(1) AS n, MAX(timestamp) AS last_ts - FROM tool_log - WHERE tool_name LIKE 'mcp__%' - GROUP BY tool_name, specialist - ORDER BY n DESC, last_ts DESC - LIMIT ? - """, - (max(1, int(limit)),), - ).fetchall() + rows = self._tool_log_queries_repo().list_mcp_tool_usage_summary(limit=max(1, int(limit))) out: list[dict[str, Any]] = [] for r in rows: tool_name = str(r["tool_name"] or "") @@ -2336,15 +2068,7 @@ class SqliteStore: def list_mcp_tool_aggregate_usage(self) -> dict[str, dict[str, Any]]: """Cross-session counts and last call time per MCP tool name (``mcp__*``).""" - with self._connect() as conn: - rows = conn.execute( - """ - SELECT tool_name, COUNT(1) AS n, MAX(timestamp) AS last_ts - FROM tool_log - WHERE tool_name LIKE 'mcp__%' - GROUP BY tool_name - """ - ).fetchall() + rows = self._tool_log_queries_repo().list_mcp_tool_aggregate_usage() out: dict[str, dict[str, Any]] = {} for r in rows: tn = str(r["tool_name"] or "") @@ -2354,24 +2078,9 @@ class SqliteStore: return out def list_mcp_tool_call_logs(self, *, server_id: str | None = None, limit: int = 200) -> list[dict[str, Any]]: - where = "WHERE tool_name LIKE 'mcp__%'" - params: list[Any] = [] - sid = str(server_id or "").strip() - if sid: - where += " AND tool_name LIKE ?" - params.append(f"mcp__{sid}__%") - params.append(max(1, int(limit))) - with self._connect() as conn: - rows = conn.execute( - f""" - SELECT session_id, tool_name, specialist, args, result, timestamp, duration_ms - FROM tool_log - {where} - ORDER BY id DESC - LIMIT ? - """, - tuple(params), - ).fetchall() + rows = self._tool_log_queries_repo().list_mcp_tool_call_logs( + server_id=server_id, limit=max(1, int(limit)) + ) out: list[dict[str, Any]] = [] for r in rows: try: @@ -2418,42 +2127,13 @@ class SqliteStore: } def count_messages(self, session_id: str) -> int: - with self._connect() as conn: - row = conn.execute( - "SELECT COUNT(*) AS c FROM chat_message WHERE session_id = ?", - (session_id,), - ).fetchone() - return int(row["c"]) if row else 0 + return self._chat_messages_repo().count_messages(session_id=session_id) def get_session_messages_meta(self, session_id: str) -> SessionMessagesMeta: - with self._connect() as conn: - row = conn.execute( - """ - SELECT - COUNT(*) AS c, - MAX(id) AS last_id, - MAX(timestamp) AS last_ts - FROM chat_message - WHERE session_id = ? - """, - (session_id,), - ).fetchone() - return SessionMessagesMeta( - session_id=session_id, - message_count=int(row["c"] or 0) if row else 0, - last_message_id=int(row["last_id"]) if row and row["last_id"] is not None else None, - last_message_at=str(row["last_ts"]) if row and row["last_ts"] is not None else None, - ) + return self._chat_messages_repo().session_messages_meta(session_id=session_id) def get_last_message_id(self, session_id: str) -> int | None: - with self._connect() as conn: - row = conn.execute( - "SELECT MAX(id) AS m FROM chat_message WHERE session_id = ?", - (session_id,), - ).fetchone() - if not row or row["m"] is None: - return None - return int(row["m"]) + return self._chat_messages_repo().last_message_id(session_id=session_id) def ensure_default_session(self) -> ChatSession: sessions = self.list_sessions(limit=1, offset=0) @@ -2464,157 +2144,60 @@ class SqliteStore: return self.create_session(title) def rename_session(self, session_id: str, title: str) -> None: - with self._connect() as conn: - conn.execute( - "UPDATE chat_session SET title = ? WHERE id = ?", - (title, session_id), - ) + self._chat_sessions_repo().rename_chat_session(session_id=session_id, title=title) def trim_messages(self, session_id: str, keep_last: int) -> None: if keep_last <= 0: self.delete_session(session_id) return - with self._connect() as conn: - rows = list( - conn.execute( - """ - SELECT id, role, tool_calls FROM chat_message - WHERE session_id = ? ORDER BY id ASC - """, - (session_id,), - ) - ) - start = _trim_messages_start_index(rows, keep_last) - if start is None: - return - min_keep_id = int(rows[start]["id"]) - with self._connect() as conn: - conn.execute( - "DELETE FROM chat_message WHERE session_id = ? AND id < ?", - (session_id, min_keep_id), - ) + self._chat_messages_repo().trim_messages_keep_last(session_id=session_id, keep_last=keep_last) def fork_session(self, source_session_id: str, up_to_message_id: int, title: str) -> ChatSession: """将 `id <= up_to_message_id` 的消息复制到新会话,并重映射 tool 的 assistant_message_id。""" - with self._connect() as conn: - chk = conn.execute( - "SELECT 1 FROM chat_message WHERE session_id = ? AND id = ? LIMIT 1", - (source_session_id, up_to_message_id), - ).fetchone() - if not chk: - raise ValueError("message not in session") - rows = list( - conn.execute( - """ - SELECT id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - FROM chat_message - WHERE session_id = ? AND id <= ? - ORDER BY id ASC - """, - (source_session_id, up_to_message_id), - ) - ) + self._chat_messages_repo().fork_assert_anchor( + source_session_id=source_session_id, + up_to_message_id=up_to_message_id, + ) new_sess = self.create_session(title) - id_map: dict[int, int] = {} - with self._connect() as conn: - for r in rows: - old_id = int(r["id"]) - role = str(r["role"]) - tool_calls_text = r["tool_calls"] - if role == "tool" and tool_calls_text: - try: - meta = json.loads(tool_calls_text) - if isinstance(meta, dict): - aid = meta.get("assistant_message_id") - if aid is not None: - new_aid = id_map.get(int(aid)) - if new_aid is not None: - meta = {**meta, "assistant_message_id": new_aid} - tool_calls_text = json.dumps(meta, ensure_ascii=False) - except (json.JSONDecodeError, TypeError, ValueError): - pass - cur = conn.execute( - """ - INSERT INTO chat_message ( - session_id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - ) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - new_sess.id, - role, - r["content"], - tool_calls_text, - r["attachments"], - r["turn_uuid"], - r["event_type"], - r["event_payload"], - r["timestamp"], - ), - ) - id_map[old_id] = int(cur.lastrowid) - last_ts = rows[-1]["timestamp"] if rows else utc_now_iso() - conn.execute( - "UPDATE chat_session SET last_message_at = ? WHERE id = ?", - (last_ts, new_sess.id), - ) + self._chat_messages_repo().fork_copy_messages_to_session( + source_session_id=source_session_id, + up_to_message_id=up_to_message_id, + new_session_id=new_sess.id, + ) return self.get_session(new_sess.id) or new_sess def set_setting(self, key: str, value: str) -> None: ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - """ - INSERT INTO app_setting (key, value, is_secret, updated_at) - VALUES (?, ?, 0, ?) - ON CONFLICT(key) DO UPDATE SET value = excluded.value, is_secret = 0, updated_at = excluded.updated_at - """, - (key, value, ts), - ) + self._app_settings_repo().upsert_plain(key=key, value=value, updated_at=ts) def get_setting(self, key: str) -> Optional[str]: - with self._connect() as conn: - row = conn.execute( - "SELECT value, is_secret FROM app_setting WHERE key = ?", - (key,), - ).fetchone() - if not row: + row = self._app_settings_repo().fetch_row(key=key) + if row is None: return None - if int(row["is_secret"]) != 0: + val, is_secret = row + if is_secret != 0: return None - return str(row["value"]) + return str(val) def set_secret(self, key: str, plain_text: str) -> None: ts = utc_now_iso() enc = _encode_secret(plain_text) - with self._connect() as conn: - conn.execute( - """ - INSERT INTO app_setting (key, value, is_secret, updated_at) - VALUES (?, ?, 1, ?) - ON CONFLICT(key) DO UPDATE SET value = excluded.value, is_secret = 1, updated_at = excluded.updated_at - """, - (key, enc, ts), - ) + self._app_settings_repo().upsert_secret(key=key, encoded_value=enc, updated_at=ts) def get_secret(self, key: str) -> Optional[str]: - with self._connect() as conn: - row = conn.execute( - "SELECT value, is_secret FROM app_setting WHERE key = ?", - (key,), - ).fetchone() - if not row: + row = self._app_settings_repo().fetch_row(key=key) + if row is None: return None - if int(row["is_secret"]) != 1: + val, is_secret = row + if is_secret != 1: return None try: - return _decode_secret(str(row["value"])) + return _decode_secret(str(val)) except Exception: return None def delete_setting(self, key: str) -> None: - with self._connect() as conn: - conn.execute("DELETE FROM app_setting WHERE key = ?", (key,)) + self._app_settings_repo().delete_key(key=key) def migrate_secrets_to_fernet(self) -> dict[str, int]: """Migrate legacy b64 secrets to a stronger scheme for both app settings and llm profiles. @@ -2627,28 +2210,15 @@ class SqliteStore: if sys.platform != "win32" and _fernet() is None: raise _CryptoError("fernet is not available; set AIA_ASSISTANT_MASTER_KEY and install cryptography") - migrated_app_settings = 0 - migrated_llm_profiles = 0 ts = utc_now_iso() + migrated_app_settings = self._app_settings_repo().migrate_b64_secrets( + ts=ts, + decode_secret=_decode_secret, + encode_secret=_encode_secret, + predicate_new_encoding=lambda enc: enc.startswith("fernet:") or enc.startswith("dpapi:"), + ) + migrated_llm_profiles = 0 with self._connect() as conn: - rows = conn.execute( - "SELECT key, value FROM app_setting WHERE is_secret = 1 AND value LIKE 'b64:%'" - ).fetchall() - for r in rows: - k = str(r["key"] or "") - v = str(r["value"] or "") - try: - plain = _decode_secret(v) - except Exception: - continue - enc = _encode_secret(plain) - if enc != v and (enc.startswith("fernet:") or enc.startswith("dpapi:")): - conn.execute( - "UPDATE app_setting SET value = ?, updated_at = ? WHERE key = ? AND is_secret = 1", - (enc, ts, k), - ) - migrated_app_settings += 1 - profs = conn.execute( "SELECT id, api_key FROM llm_profile WHERE api_key IS NOT NULL AND api_key LIKE 'b64:%'" ).fetchall() @@ -2676,15 +2246,13 @@ class SqliteStore: def legacy_secret_stats(self) -> dict[str, Any]: """Return counts of legacy b64 secrets for UI warning.""" + legacy_b64_app = self._app_settings_repo().count_legacy_b64_secrets() with self._connect() as conn: - row1 = conn.execute( - "SELECT COUNT(1) AS n FROM app_setting WHERE is_secret = 1 AND value LIKE 'b64:%'" - ).fetchone() row2 = conn.execute( "SELECT COUNT(1) AS n FROM llm_profile WHERE api_key IS NOT NULL AND api_key LIKE 'b64:%'" ).fetchone() return { - "legacy_b64_app_settings": int((row1["n"] if row1 else 0) or 0), + "legacy_b64_app_settings": int(legacy_b64_app), "legacy_b64_llm_profiles": int((row2["n"] if row2 else 0) or 0), } @@ -4106,27 +3674,19 @@ class SqliteStore: payload: dict[str, Any], ) -> None: ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - """ - INSERT INTO trace_event - (session_id, trace_id, span_id, parent_span_id, event_type, payload, timestamp) - VALUES (?, ?, ?, ?, ?, ?, ?) - """, - ( - str(session_id), - str(trace_id), - str(span_id), - str(parent_span_id) if parent_span_id else None, - str(event_type), - json.dumps(payload or {}, ensure_ascii=False, default=str), - ts, - ), - ) + self._trace_events_repo().insert_one( + session_id=str(session_id), + trace_id=str(trace_id), + span_id=str(span_id), + parent_span_id=str(parent_span_id) if parent_span_id else None, + event_type=str(event_type), + payload=json.dumps(payload or {}, ensure_ascii=False, default=str), + timestamp=ts, + ) def add_trace_events_batch(self, events: list[dict[str, Any]]) -> None: - rows: list[tuple[str, str, str, str | None, str, str, str]] = [] ts = utc_now_iso() + batch: list[dict[str, Any]] = [] for e in events or []: try: session_id = str(e.get("session_id") or "").strip() @@ -4137,43 +3697,25 @@ class SqliteStore: parent_span_id = str(e.get("parent_span_id") or "").strip() or None event_type = str(e.get("event_type") or "").strip() payload = e.get("payload") if isinstance(e.get("payload"), dict) else {} - rows.append( - ( - session_id, - trace_id, - span_id, - parent_span_id, - event_type, - json.dumps(payload or {}, ensure_ascii=False, default=str), - ts, - ) + batch.append( + { + "session_id": session_id, + "trace_id": trace_id, + "span_id": span_id, + "parent_span_id": parent_span_id, + "event_type": event_type, + "payload": json.dumps(payload or {}, ensure_ascii=False, default=str), + "timestamp": ts, + } ) except Exception: continue - if not rows: - return - with self._connect() as conn: - conn.executemany( - """ - INSERT INTO trace_event - (session_id, trace_id, span_id, parent_span_id, event_type, payload, timestamp) - VALUES (?, ?, ?, ?, ?, ?, ?) - """, - rows, - ) + self._trace_events_repo().insert_many(batch) def list_trace_events(self, *, session_id: str, limit: int = 300) -> list[dict[str, Any]]: - with self._connect() as conn: - rows = conn.execute( - """ - SELECT trace_id, span_id, parent_span_id, event_type, payload, timestamp - FROM trace_event - WHERE session_id = ? - ORDER BY id DESC - LIMIT ? - """, - (str(session_id), max(1, int(limit))), - ).fetchall() + rows = self._trace_events_repo().list_trace_events_desc( + session_id=str(session_id), limit=max(1, int(limit)) + ) out: list[dict[str, Any]] = [] for r in rows: try: @@ -4195,21 +3737,9 @@ class SqliteStore: def list_trace_events_for_trace( self, *, session_id: str, trace_id: str, limit: int = 500 ) -> list[dict[str, Any]]: - sid = str(session_id or "").strip() - tid = str(trace_id or "").strip() - if not sid or not tid: - return [] - with self._connect() as conn: - rows = conn.execute( - """ - SELECT trace_id, span_id, parent_span_id, event_type, payload, timestamp - FROM trace_event - WHERE session_id = ? AND trace_id = ? - ORDER BY id ASC - LIMIT ? - """, - (sid, tid, max(1, int(limit))), - ).fetchall() + rows = self._trace_events_repo().list_trace_events_for_trace_asc( + session_id=session_id, trace_id=trace_id, limit=max(1, int(limit)) + ) out: list[dict[str, Any]] = [] for r in rows: try: @@ -4234,16 +3764,9 @@ class SqliteStore: tid = str(trace_id or "").strip() if not sid or not tid: return None, None - with self._connect() as conn: - rows = conn.execute( - """ - SELECT event_type, timestamp - FROM trace_event - WHERE session_id = ? AND trace_id = ? - ORDER BY id ASC - """, - (sid, tid), - ).fetchall() + rows = self._trace_events_repo().list_event_type_timestamp_for_trace( + session_id=sid, trace_id=tid + ) if not rows: return None, None start = None @@ -4271,34 +3794,9 @@ class SqliteStore: if not start or not end: return [] lim = max(1, min(int(limit), 2000)) - with self._connect() as conn: - rows = conn.execute( - """ - SELECT id, session_id, role, content, tool_calls, attachments, turn_uuid, event_type, event_payload, timestamp - FROM chat_message - WHERE session_id = ? AND timestamp >= ? AND timestamp <= ? - ORDER BY id ASC - LIMIT ? - """, - (sid, start, end, lim), - ).fetchall() - out: list[dict[str, Any]] = [] - for r in rows: - out.append( - { - "id": int(r["id"] or 0), - "session_id": str(r["session_id"] or ""), - "role": str(r["role"] or ""), - "content": str(r["content"] or ""), - "tool_calls": r["tool_calls"], - "attachments": r["attachments"], - "turn_uuid": str(r["turn_uuid"] or ""), - "event_type": str(r["event_type"] or ""), - "event_payload": r["event_payload"], - "timestamp": str(r["timestamp"] or ""), - } - ) - return out + return self._chat_messages_repo().list_messages_in_time_window( + session_id=sid, start_ts=start, end_ts=end, limit=lim + ) # ---------------------------- # Tenant / User / Bind Codes @@ -4306,46 +3804,39 @@ class SqliteStore: def create_tenant(self, name: str) -> dict[str, Any]: tid = str(uuid.uuid4()) ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - "INSERT INTO tenant (id, name, created_at) VALUES (?, ?, ?)", - (tid, str(name or "").strip() or "Team", ts), - ) + nm = str(name or "").strip() or "Team" + self._tenant_repo().insert_tenant(tenant_id=tid, name=nm, created_at=ts) return {"id": tid, "name": name, "created_at": ts} def delete_tenant(self, *, tenant_id: str) -> int: tid = str(tenant_id or "").strip() if not tid: return 0 - with self._connect() as conn: - cur = conn.execute("DELETE FROM tenant WHERE id = ?", (tid,)) - return int(cur.rowcount or 0) + return self._tenant_repo().delete_tenant(tenant_id=tid) def list_tenants(self, *, limit: int = 200) -> list[dict[str, Any]]: - with self._connect() as conn: - rows = conn.execute( - "SELECT id, name, created_at FROM tenant ORDER BY created_at DESC LIMIT ?", - (max(1, int(limit)),), - ).fetchall() - return [{"id": r["id"], "name": r["name"], "created_at": r["created_at"]} for r in rows] + return self._tenant_repo().list_tenants(limit=max(1, int(limit))) def create_user(self, *, tenant_id: str, display_name: str, role: str) -> dict[str, Any]: uid = str(uuid.uuid4()) ts = utc_now_iso() username = (str(display_name or "").strip() or "user").lower().replace(" ", "_") - with self._connect() as conn: - # avoid conflicts within tenant. - row = conn.execute( - "SELECT COUNT(*) AS c FROM app_user WHERE tenant_id = ? AND username = ?", - (str(tenant_id), username), - ).fetchone() - suffix = int(row["c"] or 0) if row else 0 - if suffix > 0: - username = f"{username}_{suffix+1}" - conn.execute( - "INSERT INTO app_user (id, tenant_id, username, display_name, role, password_hash, is_active, created_at) VALUES (?, ?, ?, ?, ?, '', 1, ?)", - (uid, str(tenant_id), username, str(display_name or "").strip() or "User", str(role or "member"), ts), - ) + repo = self._app_users_repo() + suffix = repo.count_by_tenant_username(tenant_id=str(tenant_id), username=username) + if suffix > 0: + username = f"{username}_{suffix+1}" + disp = str(display_name or "").strip() or "User" + repo.insert_user( + user_id=uid, + tenant_id=str(tenant_id), + username=username, + display_name=disp, + role=str(role or "member"), + password_hash="", + is_active=1, + created_at=ts, + avatar_attachment_id=None, + ) return { "id": uid, "tenant_id": tenant_id, @@ -4365,174 +3856,26 @@ class SqliteStore: q: str | None = None, include_inactive: bool = True, ) -> list[dict[str, Any]]: - where = ["tenant_id = ?"] - params: list[Any] = [str(tenant_id)] - token = str(q or "").strip() - if token: - where.append("(LOWER(display_name) LIKE ? OR LOWER(COALESCE(username,'')) LIKE ? OR id LIKE ?)") - key = f"%{token.lower()}%" - params.extend([key, key, f"%{token[:32]}%"]) - if not include_inactive: - where.append("COALESCE(is_active, 1) = 1") - wsql = " AND ".join(where) - with self._connect() as conn: - rows = conn.execute( - """ - SELECT - id, - tenant_id, - username, - display_name, - role, - COALESCE(is_active,1) AS is_active, - created_at, - (CASE WHEN TRIM(COALESCE(password_hash,'')) != '' THEN 1 ELSE 0 END) AS has_password, - ( - EXISTS ( - SELECT 1 FROM channel_identity_v2 ci - WHERE ci.tenant_id = app_user.tenant_id AND ci.user_id = app_user.id AND ci.channel = 'wecom' - ) - OR EXISTS ( - SELECT 1 FROM channel_identity ci - WHERE ci.tenant_id = app_user.tenant_id AND ci.user_id = app_user.id AND ci.channel = 'wecom' - ) - ) AS wecom_linked, - ( - EXISTS ( - SELECT 1 FROM channel_identity_v2 ci - WHERE ci.tenant_id = app_user.tenant_id AND ci.user_id = app_user.id - ) - OR EXISTS ( - SELECT 1 FROM channel_identity ci - WHERE ci.tenant_id = app_user.tenant_id AND ci.user_id = app_user.id - ) - ) AS channel_linked, - ( - SELECT GROUP_CONCAT(z.eid, ', ') - FROM ( - SELECT DISTINCT TRIM(ci.external_user_id) AS eid - FROM channel_identity ci - WHERE ci.tenant_id = app_user.tenant_id - AND ci.user_id = app_user.id - AND ci.channel = 'wecom' - AND TRIM(COALESCE(ci.external_user_id, '')) != '' - UNION - SELECT DISTINCT TRIM(ci.external_user_id) AS eid - FROM channel_identity_v2 ci - WHERE ci.tenant_id = app_user.tenant_id - AND ci.user_id = app_user.id - AND ci.channel = 'wecom' - AND TRIM(COALESCE(ci.external_user_id, '')) != '' - ) AS z - ) AS wecom_external_user_ids - FROM app_user - WHERE """ - + wsql - + """ - ORDER BY created_at DESC - LIMIT ? - OFFSET ? - """, - (*params, max(1, int(limit)), max(0, int(offset))), - ).fetchall() - out: list[dict[str, Any]] = [] - for r in rows: - has_pw = bool(int(r["has_password"] or 0)) - uname = str(r["username"] or "") - can_chat = bool(has_pw) - out.append( - { - "id": r["id"], - "tenant_id": r["tenant_id"], - "username": r["username"], - "display_name": r["display_name"], - "role": r["role"], - "is_active": bool(int(r["is_active"] or 0)), - "created_at": r["created_at"], - "has_password": has_pw, - "wecom_linked": bool(int(r["wecom_linked"] or 0)), - "channel_linked": bool(int(r["channel_linked"] or 0)), - "can_chat_login": can_chat, - "wecom_external_user_ids": str(r["wecom_external_user_ids"] or "").strip(), - } - ) - return out + return self._app_users_repo().list_users_for_tenant( + tenant_id=str(tenant_id), + limit=limit, + offset=offset, + q=q, + include_inactive=include_inactive, + ) def get_user_by_username(self, *, tenant_id: str, username: str) -> dict[str, Any] | None: - with self._connect() as conn: - r = conn.execute( - """ - SELECT id, tenant_id, username, display_name, role, COALESCE(is_active,1) AS is_active, created_at, COALESCE(password_hash,'') AS password_hash, COALESCE(avatar_attachment_id,'') AS avatar_attachment_id - FROM app_user - WHERE tenant_id = ? AND username = ? - LIMIT 1 - """, - (str(tenant_id), str(username)), - ).fetchone() - if not r: - return None - return { - "id": r["id"], - "tenant_id": r["tenant_id"], - "username": r["username"], - "display_name": r["display_name"], - "role": r["role"], - "is_active": bool(int(r["is_active"] or 0)), - "created_at": r["created_at"], - "password_hash": r["password_hash"], - "avatar_attachment_id": str(r["avatar_attachment_id"] or "").strip() or None, - } + return self._app_users_repo().fetch_by_tenant_and_username( + tenant_id=str(tenant_id), username=str(username) + ) def get_user_by_username_global(self, *, username: str) -> dict[str, Any] | None: - with self._connect() as conn: - r = conn.execute( - """ - SELECT id, tenant_id, username, display_name, role, COALESCE(is_active,1) AS is_active, created_at, COALESCE(password_hash,'') AS password_hash, COALESCE(avatar_attachment_id,'') AS avatar_attachment_id - FROM app_user - WHERE username = ? - ORDER BY created_at ASC - LIMIT 1 - """, - (str(username),), - ).fetchone() - if not r: - return None - return { - "id": r["id"], - "tenant_id": r["tenant_id"], - "username": r["username"], - "display_name": r["display_name"], - "role": r["role"], - "is_active": bool(int(r["is_active"] or 0)), - "created_at": r["created_at"], - "password_hash": r["password_hash"], - "avatar_attachment_id": str(r["avatar_attachment_id"] or "").strip() or None, - } + return self._app_users_repo().fetch_first_by_username_global(username=str(username)) def get_user_by_id(self, *, tenant_id: str, user_id: str) -> dict[str, Any] | None: - with self._connect() as conn: - r = conn.execute( - """ - SELECT id, tenant_id, username, display_name, role, COALESCE(is_active,1) AS is_active, created_at, COALESCE(password_hash,'') AS password_hash, COALESCE(avatar_attachment_id,'') AS avatar_attachment_id - FROM app_user - WHERE tenant_id = ? AND id = ? - LIMIT 1 - """, - (str(tenant_id), str(user_id)), - ).fetchone() - if not r: - return None - return { - "id": r["id"], - "tenant_id": r["tenant_id"], - "username": r["username"], - "display_name": r["display_name"], - "role": r["role"], - "is_active": bool(int(r["is_active"] or 0)), - "created_at": r["created_at"], - "password_hash": r["password_hash"], - "avatar_attachment_id": str(r["avatar_attachment_id"] or "").strip() or None, - } + return self._app_users_repo().fetch_by_tenant_and_id( + tenant_id=str(tenant_id), user_id=str(user_id) + ) def create_user_account( self, @@ -4546,23 +3889,17 @@ class SqliteStore: ) -> dict[str, Any]: uid = str(uuid.uuid4()) ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - """ - INSERT INTO app_user (id, tenant_id, username, display_name, role, password_hash, is_active, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?) - """, - ( - uid, - str(tenant_id), - str(username).strip(), - str(display_name).strip() or str(username).strip(), - str(role or "member"), - str(password_hash or ""), - 1 if is_active else 0, - ts, - ), - ) + self._app_users_repo().insert_user( + user_id=uid, + tenant_id=str(tenant_id), + username=str(username).strip(), + display_name=str(display_name).strip() or str(username).strip(), + role=str(role or "member"), + password_hash=str(password_hash or ""), + is_active=1 if is_active else 0, + created_at=ts, + avatar_attachment_id=None, + ) return self.get_user_by_id(tenant_id=tenant_id, user_id=uid) or {} def update_user_account( @@ -4576,37 +3913,18 @@ class SqliteStore: password_hash: str | None = None, avatar_attachment_id: str | None = None, ) -> bool: - sets: list[str] = [] - params: list[Any] = [] - if display_name is not None: - sets.append("display_name = ?") - params.append(str(display_name).strip() or "User") - if role is not None: - sets.append("role = ?") - params.append(str(role).strip() or "member") - if is_active is not None: - sets.append("is_active = ?") - params.append(1 if is_active else 0) - if password_hash is not None: - sets.append("password_hash = ?") - params.append(str(password_hash)) - if avatar_attachment_id is not None: - sets.append("avatar_attachment_id = ?") - aid = str(avatar_attachment_id).strip() - params.append(aid if aid else None) - if not sets: - return False - with self._connect() as conn: - cur = conn.execute( - f"UPDATE app_user SET {', '.join(sets)} WHERE tenant_id = ? AND id = ?", - (*params, str(tenant_id), str(user_id)), - ) - return bool(cur.rowcount and cur.rowcount > 0) + return self._app_users_repo().update_user_account( + tenant_id=str(tenant_id), + user_id=str(user_id), + display_name=display_name, + role=role, + is_active=is_active, + password_hash=password_hash, + avatar_attachment_id=avatar_attachment_id, + ) def delete_user_account(self, *, tenant_id: str, user_id: str) -> int: - with self._connect() as conn: - cur = conn.execute("DELETE FROM app_user WHERE tenant_id = ? AND id = ?", (str(tenant_id), str(user_id))) - return int(cur.rowcount or 0) + return self._app_users_repo().delete_user_account(tenant_id=str(tenant_id), user_id=str(user_id)) def upsert_user_permission(self, *, tenant_id: str, user_id: str, permission: str) -> None: ts = utc_now_iso() @@ -4744,72 +4062,30 @@ class SqliteStore: expires_at: str, ) -> None: ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - """ - INSERT INTO auth_session (session_token_hash, tenant_id, user_id, role, created_at, expires_at, last_seen_at, revoked_at) - VALUES (?, ?, ?, ?, ?, ?, ?, NULL) - """, - (str(session_token_hash), str(tenant_id), str(user_id), str(role), ts, str(expires_at), ts), - ) + self._auth_sessions_repo().insert_session( + session_token_hash=str(session_token_hash), + tenant_id=str(tenant_id), + user_id=str(user_id), + role=str(role), + created_at=ts, + expires_at=str(expires_at), + last_seen_at=ts, + ) def revoke_auth_session(self, *, session_token_hash: str) -> int: ts = utc_now_iso() - with self._connect() as conn: - cur = conn.execute( - """ - UPDATE auth_session - SET revoked_at = ? - WHERE session_token_hash = ? AND revoked_at IS NULL - """, - (ts, str(session_token_hash)), - ) - return int(cur.rowcount or 0) + return self._auth_sessions_repo().revoke_one(session_token_hash=str(session_token_hash), revoked_at=ts) def revoke_all_auth_sessions(self) -> int: ts = utc_now_iso() - with self._connect() as conn: - cur = conn.execute( - """ - UPDATE auth_session - SET revoked_at = ? - WHERE revoked_at IS NULL - """, - (ts,), - ) - return int(cur.rowcount or 0) + return self._auth_sessions_repo().revoke_all_active(revoked_at=ts) def get_auth_session(self, *, session_token_hash: str) -> dict[str, Any] | None: - with self._connect() as conn: - r = conn.execute( - """ - SELECT session_token_hash, tenant_id, user_id, role, created_at, expires_at, last_seen_at, revoked_at - FROM auth_session - WHERE session_token_hash = ? - LIMIT 1 - """, - (str(session_token_hash),), - ).fetchone() - if not r: - return None - return { - "session_token_hash": r["session_token_hash"], - "tenant_id": r["tenant_id"], - "user_id": r["user_id"], - "role": r["role"], - "created_at": r["created_at"], - "expires_at": r["expires_at"], - "last_seen_at": r["last_seen_at"], - "revoked_at": r["revoked_at"], - } + return self._auth_sessions_repo().fetch_by_hash(session_token_hash=str(session_token_hash)) def touch_auth_session(self, *, session_token_hash: str) -> None: ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - "UPDATE auth_session SET last_seen_at = ? WHERE session_token_hash = ?", - (ts, str(session_token_hash)), - ) + self._auth_sessions_repo().touch(session_token_hash=str(session_token_hash), last_seen_at=ts) def add_admin_audit_log( self, @@ -5161,14 +4437,9 @@ class SqliteStore: def create_bind_code(self, *, tenant_id: str, role: str, code: str) -> dict[str, Any]: ts = utc_now_iso() - with self._connect() as conn: - conn.execute( - """ - INSERT INTO bind_code (code, tenant_id, role, created_at, used_at, used_by_external_user_id) - VALUES (?, ?, ?, ?, NULL, NULL) - """, - (str(code), str(tenant_id), str(role), ts), - ) + self._bind_code_repo().insert_bind_code( + code=str(code), tenant_id=str(tenant_id), role=str(role), created_at=ts + ) return {"code": code, "tenant_id": tenant_id, "role": role, "created_at": ts} def consume_bind_code( @@ -5178,61 +4449,27 @@ class SqliteStore: code = str(code or "").strip() if not code: return None - with self._connect() as conn: - row = conn.execute( - """ - SELECT code, tenant_id, role, used_at - FROM bind_code - WHERE code = ? - """, - (code,), - ).fetchone() - if not row or row["used_at"]: - return None - tenant_id = str(row["tenant_id"]) - role = str(row["role"] or "member") - user = self.create_user(tenant_id=tenant_id, display_name=display_name or "User", role=role) - self.upsert_channel_identity( - tenant_id=tenant_id, - channel=channel, - external_user_id=external_user_id, - user_id=str(user["id"]), - ) - ts = utc_now_iso() - conn.execute( - """ - UPDATE bind_code - SET used_at = ?, used_by_external_user_id = ? - WHERE code = ? - """, - (ts, str(external_user_id), code), - ) + row = self._bind_code_repo().fetch_by_code(code=code) + if not row or row.get("used_at"): + return None + tenant_id = str(row["tenant_id"]) + role = str(row["role"] or "member") + user = self.create_user(tenant_id=tenant_id, display_name=display_name or "User", role=role) + self.upsert_channel_identity( + tenant_id=tenant_id, + channel=channel, + external_user_id=external_user_id, + user_id=str(user["id"]), + ) + ts = utc_now_iso() + self._bind_code_repo().mark_used( + code=code, used_at=ts, used_by_external_user_id=str(external_user_id) + ) return {"tenant_id": tenant_id, "user_id": user["id"], "role": role} def list_bind_codes(self, *, tenant_id: str | None = None, limit: int = 200) -> list[dict[str, Any]]: lim = max(1, int(limit)) - with self._connect() as conn: - if tenant_id: - rows = conn.execute( - """ - SELECT code, tenant_id, role, created_at, used_at, used_by_external_user_id - FROM bind_code - WHERE tenant_id = ? - ORDER BY created_at DESC - LIMIT ? - """, - (str(tenant_id), lim), - ).fetchall() - else: - rows = conn.execute( - """ - SELECT code, tenant_id, role, created_at, used_at, used_by_external_user_id - FROM bind_code - ORDER BY created_at DESC - LIMIT ? - """, - (lim,), - ).fetchall() + rows = self._bind_code_repo().list_bind_codes(tenant_id=tenant_id, limit=lim) return [ { "code": r["code"], @@ -5414,22 +4651,7 @@ class SqliteStore: return str(sess.id) def backfill_ui_session_owner_from_channel_v2(self) -> int: - with self._connect() as conn: - cur = conn.execute( - """ - INSERT OR IGNORE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) - SELECT DISTINCT cs.session_id, cs.tenant_id, ci.user_id, ? - FROM channel_session_v2 cs - JOIN channel_identity_v2 ci - ON ci.tenant_id = cs.tenant_id - AND ci.channel = cs.channel - AND ci.account_id = cs.account_id - AND ci.external_user_id = cs.external_user_id - WHERE cs.session_id IS NOT NULL AND cs.session_id != '' - """, - (utc_now_iso(),), - ) - return int(cur.rowcount or 0) + return self._ui_session_owner_repo().backfill_from_channel_v2(created_at=utc_now_iso()) # ---------------------------- # Oclaw tasks @@ -5682,6 +4904,26 @@ class SqliteStore: ) -> int: ts = utc_now_iso() with self._connect() as conn: + if self._use_pg: + cur = conn.execute( + """ + INSERT INTO oclaw_attempt(run_id, tenant_id, session_id, attempt_no, status, reason, payload, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + RETURNING id + """, + ( + str(run_id), + str(tenant_id), + str(session_id), + int(attempt_no), + str(status), + str(reason or ""), + json.dumps(payload or {}, ensure_ascii=False), + ts, + ), + ) + row = cur.fetchone() + return int(row["id"]) if row else 0 cur = conn.execute( """ INSERT INTO oclaw_attempt(run_id, tenant_id, session_id, attempt_no, status, reason, payload, created_at) diff --git a/tests/test_admin_auth_rbac.py b/tests/test_admin_auth_rbac.py index f96eefe9..d1204547 100644 --- a/tests/test_admin_auth_rbac.py +++ b/tests/test_admin_auth_rbac.py @@ -8,8 +8,9 @@ from pathlib import Path from fastapi.testclient import TestClient from interfaces.http.fastapi_app import create_app -from svc.config.paths import db_path +from svc.persistence.db.engine import clear_assistant_engine_cache from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store class AdminAuthRBACTests(unittest.TestCase): @@ -36,6 +37,7 @@ class AdminAuthRBACTests(unittest.TestCase): self.client = TestClient(app) def tearDown(self) -> None: + clear_assistant_engine_cache() self._tmp.cleanup() def _login(self, username: str = "administrator", password: str = "test-admin-pass") -> str: @@ -65,7 +67,7 @@ class AdminAuthRBACTests(unittest.TestCase): def test_cross_tenant_forbidden(self) -> None: token = self._login() - other = SqliteStore(db_path()).create_tenant("Other") + other = get_assistant_store().create_tenant("Other") r = self.client.get( f"/admin/api/users?tenant_id={other['id']}", headers={"authorization": f"Bearer {token}"}, @@ -74,7 +76,7 @@ class AdminAuthRBACTests(unittest.TestCase): self.assertIn(r.status_code, (200, 403)) def test_chat_login_allows_non_administrator_with_password(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.create_user_account( tenant_id=self.tenant_id, username="alice", @@ -98,7 +100,7 @@ class AdminAuthRBACTests(unittest.TestCase): self.assertEqual(str((data.get("session") or {}).get("username") or ""), "alice") def test_chat_login_username_case_insensitive(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.create_user_account( tenant_id=self.tenant_id, username="carol", @@ -122,7 +124,7 @@ class AdminAuthRBACTests(unittest.TestCase): self.assertEqual(str((data.get("session") or {}).get("username") or ""), "carol") def test_console_login_allows_member_username_with_admin_read(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.create_user_account( tenant_id=self.tenant_id, username="bob", diff --git a/tests/test_chat_api_message_filter.py b/tests/test_chat_api_message_filter.py index 590c047c..de4039cb 100644 --- a/tests/test_chat_api_message_filter.py +++ b/tests/test_chat_api_message_filter.py @@ -29,7 +29,9 @@ def test_filter_internal_instruction_user_messages_hides_from_assistant_event_pa role="assistant", event_type="assistant_text", content="已修复", - event_payload={"reasoning_content": f"任务分配\nspecialist=ops\ninstruction:\n{polluted_text}"}, + event_payload={ + "reasoning_content": f"任务分配\nspecialist=ops\ninstruction:\n{polluted_text}" + }, ), ] out = _filter_internal_instruction_user_messages(rows) @@ -45,3 +47,35 @@ def test_filter_internal_instruction_user_messages_keeps_normal_user_rows() -> N out = _filter_internal_instruction_user_messages(rows) assert len(out) == 2 assert str(getattr(out[0], "content", "")) == "你好" + + +def test_filter_hides_user_when_en_task_assignment_block() -> None: + polluted_text = "Fix the gateway." + rows = [ + SimpleNamespace(role="user", event_type="user_text", content=polluted_text), + SimpleNamespace( + role="assistant", + event_type="reasoning", + content=f"Task assignment\nspecialist=ops\ninstruction:\n{polluted_text}", + ), + SimpleNamespace(role="assistant", event_type="assistant_text", content="done"), + ] + out = _filter_internal_instruction_user_messages(rows) + assert len(out) == 2 + assert [str(getattr(x, "event_type", "")) for x in out] == ["reasoning", "assistant_text"] + + +def test_filter_keeps_user_when_reasoning_echoes_instruction_without_dispatch_header() -> None: + """Thinking traces may contain ``instruction:\\n`` + the user's words without a 任务分配 block.""" + rows = [ + SimpleNamespace(role="user", event_type="user_text", content="几点了"), + SimpleNamespace( + role="assistant", + event_type="assistant_text", + content="", + event_payload={"reasoning_content": "用户问几点了。\ninstruction:\n几点了"}, + ), + ] + out = _filter_internal_instruction_user_messages(rows) + assert len(out) == 2 + assert str(getattr(out[0], "content", "")) == "几点了" diff --git a/tests/test_chat_session_full_smoke.py b/tests/test_chat_session_full_smoke.py index e819f8ae..5cbad299 100644 --- a/tests/test_chat_session_full_smoke.py +++ b/tests/test_chat_session_full_smoke.py @@ -18,9 +18,24 @@ class ChatSessionFullSmokeTests(unittest.TestCase): def setUp(self) -> None: self._tmp = tempfile.TemporaryDirectory(ignore_cleanup_errors=True) self.db = Path(self._tmp.name) / "ops.sqlite" + # 本机若已 export PG 助手库,get_assistant_store 会连 PG,与本测试临时 SQLite 种子不一致 → login 失败。 + for k in ( + "AIA_ASSISTANT_DB_BACKEND", + "OPS_ASSISTANT_DB_BACKEND", + "AIA_ASSISTANT_DATABASE_URL", + "OPS_ASSISTANT_DATABASE_URL", + "AIA_ASSISTANT_PG_DSN", + "OPS_ASSISTANT_PG_DSN", + ): + os.environ.pop(k, None) os.environ["OPS_ASSISTANT_DB_PATH"] = str(self.db) os.environ["OPS_ASSISTANT_PASSWORD"] = "test-admin-pass" os.environ["OPS_ASSISTANT_MODE"] = "rule" + from svc.persistence.assistant_store import reset_assistant_store_singleton + from svc.persistence.db.engine import clear_assistant_engine_cache + + clear_assistant_engine_cache() + reset_assistant_store_singleton() store = SqliteStore(str(self.db)) t = store.create_tenant("Team") self.tenant_id = str(t["id"]) diff --git a/tests/test_database_backend.py b/tests/test_database_backend.py new file mode 100644 index 00000000..7d99b16f --- /dev/null +++ b/tests/test_database_backend.py @@ -0,0 +1,84 @@ +"""Tests for assistant DB backend selection.""" + +from __future__ import annotations + +import os + +import pytest + +from svc.config import database as db_cfg +from svc.persistence import assistant_store as as_mod + + +def test_assistant_db_backend_default(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.delenv("OPS_ASSISTANT_DB_BACKEND", raising=False) + assert db_cfg.assistant_db_backend() == "sqlite" + + +def test_assistant_db_backend_aliases(monkeypatch: pytest.MonkeyPatch) -> None: + for raw in ("postgresql", "POSTGRES", "pg"): + monkeypatch.setenv("AIA_ASSISTANT_DB_BACKEND", raw) + assert db_cfg.assistant_db_backend() == "postgresql" + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.delenv("OPS_ASSISTANT_DB_BACKEND", raising=False) + + +def test_assistant_postgres_dsn_required(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AIA_ASSISTANT_DB_BACKEND", "postgresql") + monkeypatch.delenv("AIA_ASSISTANT_DATABASE_URL", raising=False) + monkeypatch.delenv("OPS_ASSISTANT_DATABASE_URL", raising=False) + monkeypatch.delenv("AIA_ASSISTANT_PG_DSN", raising=False) + monkeypatch.delenv("OPS_ASSISTANT_PG_DSN", raising=False) + with pytest.raises(ValueError, match="DATABASE_URL"): + db_cfg.assistant_postgres_dsn() + + +def test_get_assistant_store_sqlite(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(tmp_path / "t.sqlite")) + as_mod.reset_assistant_store_singleton() + from svc.persistence.sqlite_store import SqliteStore + + s = as_mod.get_assistant_store() + assert isinstance(s, SqliteStore) + assert as_mod.get_assistant_store() is s + + +def test_assistant_sqlalchemy_url_sqlite(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(tmp_path / "t.sqlite")) + u = db_cfg.assistant_sqlalchemy_url() + assert u.startswith("sqlite+pysqlite:///") + + +@pytest.mark.skipif( + not (os.getenv("AIA_TEST_PG_URL") or "").strip(), + reason="Set AIA_TEST_PG_URL to run PostgreSQL integration test", +) +def test_get_assistant_store_postgresql_smoke(monkeypatch: pytest.MonkeyPatch) -> None: + as_mod.reset_assistant_store_singleton() + monkeypatch.setenv("AIA_ASSISTANT_DB_BACKEND", "postgresql") + monkeypatch.setenv("AIA_ASSISTANT_DATABASE_URL", os.environ["AIA_TEST_PG_URL"].strip()) + from svc.persistence.sqlite_store import SqliteStore + + s = as_mod.get_assistant_store() + assert isinstance(s, SqliteStore) + assert s._use_pg is True + assert s.get_setting("no_such_key_integration_test") is None + sess = s.create_session("pg_smoke_session") + assert sess.id + msg = s.add_message(sess.id, "user", "pg-smoke-body") + assert msg.id > 0 + msgs = s.get_messages(sess.id, limit=10) + assert len(msgs) == 1 + assert msgs[0].content == "pg-smoke-body" + + +def test_assistant_db_backend_invalid(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AIA_ASSISTANT_DB_BACKEND", "mysql") + with pytest.raises(ValueError, match="Invalid assistant DB backend"): + db_cfg.assistant_db_backend() + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) diff --git a/tests/test_fastapi_startup_prebuild.py b/tests/test_fastapi_startup_prebuild.py index af4e0356..ccb4e93d 100644 --- a/tests/test_fastapi_startup_prebuild.py +++ b/tests/test_fastapi_startup_prebuild.py @@ -10,8 +10,7 @@ def test_run_startup_hooks_runs_prebuild_warmup(monkeypatch) -> None: def revoke_all_auth_sessions(self) -> int: return 0 - monkeypatch.setattr(app_mod, "SqliteStore", lambda _p: DummyStore()) - monkeypatch.setattr(app_mod, "db_path", lambda: "dummy.sqlite") + monkeypatch.setattr(app_mod, "get_assistant_store", lambda: DummyStore()) monkeypatch.setattr(app_mod, "prepare_gateway_plugin_bootstrap", lambda **kwargs: {"ok": True}) monkeypatch.setattr(app_mod, "resolve_runtime_config", lambda: {}) monkeypatch.setattr(app_mod, "skill_runtime_diagnostics", lambda: {"skills_root": "/tmp", "skills_total": 0}) diff --git a/tests/test_mcp_admin_api.py b/tests/test_mcp_admin_api.py index 25f074da..cb5eb88e 100644 --- a/tests/test_mcp_admin_api.py +++ b/tests/test_mcp_admin_api.py @@ -13,6 +13,7 @@ from interfaces.http.fastapi_app import create_app from interfaces.admin import routes as admin_routes from svc.config.paths import db_path from svc.persistence.sqlite_store import SqliteStore +from svc.persistence.assistant_store import get_assistant_store class McpAdminApiTests(unittest.TestCase): @@ -106,7 +107,7 @@ class McpAdminApiTests(unittest.TestCase): self.assertEqual(toggle.status_code, 200) self.assertTrue(toggle.json().get("ok")) - store = SqliteStore(db_path()) + store = get_assistant_store() store.add_mcp_installation_log( server_id="demo-mcp", status="error", @@ -121,7 +122,7 @@ class McpAdminApiTests(unittest.TestCase): def test_healthcheck_and_tools_sync(self) -> None: script = self._write_mcp_server() - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="dummy", source_type="github", @@ -141,7 +142,7 @@ class McpAdminApiTests(unittest.TestCase): self.assertTrue(any(str(t.get("tool_name") or "") == "ping" for t in tools)) def test_healthcheck_and_tools_sync_bailian_webparser_compat(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="webparser-compat", source_type="npm", @@ -169,7 +170,7 @@ class McpAdminApiTests(unittest.TestCase): def test_reinstall_from_saved_manifest(self) -> None: script = self._write_mcp_server() - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="dummy-reinstall", source_type="npm", @@ -190,7 +191,7 @@ class McpAdminApiTests(unittest.TestCase): def test_update_from_saved_manifest_single_and_batch(self) -> None: script = self._write_mcp_server() - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="dummy-update", source_type="npm", @@ -222,7 +223,7 @@ class McpAdminApiTests(unittest.TestCase): def test_check_all_enabled_servers(self) -> None: script = self._write_mcp_server() - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="dummy-check-all", source_type="github", @@ -247,7 +248,7 @@ class McpAdminApiTests(unittest.TestCase): def test_repair_weak_skips_healthy_and_fixes_empty_tools(self) -> None: script = self._write_mcp_server() - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="dummy-repair-weak", source_type="github", @@ -286,7 +287,7 @@ class McpAdminApiTests(unittest.TestCase): def test_repair_weak_include_disabled_server(self) -> None: script = self._write_mcp_server() - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="dummy-repair-disabled", source_type="github", @@ -318,7 +319,7 @@ class McpAdminApiTests(unittest.TestCase): self.assertTrue(bool((row or {}).get("ok")), row) def test_check_updates_reports_update_candidates(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="update-s1", source_type="npm", @@ -372,7 +373,7 @@ class McpAdminApiTests(unittest.TestCase): self.assertEqual(data.get("allowed_specialists"), ["generalist", "ops"]) def test_mcp_binding_config_filters_invalid_servers(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="bind-s1", source_type="github", @@ -397,7 +398,7 @@ class McpAdminApiTests(unittest.TestCase): self.assertEqual(mapping.get("ops"), []) def test_mcp_usage_and_delete(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="usage-s1", source_type="github", @@ -431,7 +432,7 @@ class McpAdminApiTests(unittest.TestCase): self.assertEqual(left, []) def test_mcp_uninstall_remove_record(self) -> None: - store = SqliteStore(db_path()) + store = get_assistant_store() store.upsert_mcp_server( server_id="uninstall-s1", source_type="github", diff --git a/tests/test_migrate_assistant_sqlite_to_pg.py b/tests/test_migrate_assistant_sqlite_to_pg.py new file mode 100644 index 00000000..e1e74e62 --- /dev/null +++ b/tests/test_migrate_assistant_sqlite_to_pg.py @@ -0,0 +1,86 @@ +"""Tests for SQLite→PostgreSQL assistant migration ordering (no live PG required).""" + +from __future__ import annotations + +import importlib.util +import sqlite3 +import sys +from pathlib import Path + +import pytest + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +_MIGRATE_PATH = ROOT / "runtime" / "operations" / "scripts" / "migrate_assistant_sqlite_to_postgresql.py" + + +def _load_migrate_module(): + spec = importlib.util.spec_from_file_location("migrate_assistant_sqlite_to_postgresql", _MIGRATE_PATH) + assert spec and spec.loader + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + return mod + + +@pytest.fixture(scope="module") +def mig(): + return _load_migrate_module() + + +def test_migration_order_respects_foreign_keys(mig, tmp_path) -> None: + db = tmp_path / "fk.sqlite" + sl = sqlite3.connect(str(db)) + sl.executescript( + """ + CREATE TABLE tenant(id TEXT PRIMARY KEY, name TEXT NOT NULL); + CREATE TABLE app_user( + id TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL REFERENCES tenant(id), + display_name TEXT NOT NULL + ); + CREATE TABLE auth_session( + session_token_hash TEXT PRIMARY KEY, + tenant_id TEXT NOT NULL REFERENCES tenant(id), + user_id TEXT NOT NULL REFERENCES app_user(id), + role TEXT NOT NULL, + created_at TEXT NOT NULL, + expires_at TEXT NOT NULL, + last_seen_at TEXT NOT NULL + ); + """ + ) + tables = mig._sqlite_user_tables(sl) + order = mig._migration_order(sl, [t for t in tables if t in ("tenant", "app_user", "auth_session")]) + sl.close() + assert order.index("tenant") < order.index("app_user") + assert order.index("app_user") < order.index("auth_session") + + +def test_topological_sort_detects_cycle(mig) -> None: + with pytest.raises(SystemExit, match="foreign-key-safe"): + mig._topological_sort(["a", "b"], {"a": {"b"}, "b": {"a"}}) + + +def test_require_ident_rejects_injection(mig) -> None: + with pytest.raises(ValueError): + mig._require_ident("app_user;drop") + + +def test_pg_url_from_environ_order(monkeypatch: pytest.MonkeyPatch, mig) -> None: + for k in ( + "AIA_ASSISTANT_DATABASE_URL", + "OPS_ASSISTANT_DATABASE_URL", + "AIA_ASSISTANT_PG_DSN", + "OPS_ASSISTANT_PG_DSN", + ): + monkeypatch.delenv(k, raising=False) + monkeypatch.setenv("OPS_ASSISTANT_PG_DSN", "postgresql://ops/pg") + assert mig._pg_url_from_environ() == "postgresql://ops/pg" + monkeypatch.setenv("AIA_ASSISTANT_PG_DSN", "postgresql://aia/pg") + assert mig._pg_url_from_environ() == "postgresql://aia/pg" + monkeypatch.setenv("OPS_ASSISTANT_DATABASE_URL", "postgresql://ops/db") + assert mig._pg_url_from_environ() == "postgresql://ops/db" + monkeypatch.setenv("AIA_ASSISTANT_DATABASE_URL", "postgresql://aia/db") + assert mig._pg_url_from_environ() == "postgresql://aia/db" diff --git a/tests/test_persist_terminal_fallback.py b/tests/test_persist_terminal_fallback.py new file mode 100644 index 00000000..9c4735b0 --- /dev/null +++ b/tests/test_persist_terminal_fallback.py @@ -0,0 +1,64 @@ +"""persist_assistant_text_if_turn_missing""" + +from __future__ import annotations + +import uuid +from svc.persistence.sqlite_store import SqliteStore +from runtime.chat.persist_terminal_fallback import persist_assistant_text_if_turn_missing + + +def test_fallback_inserts_when_missing(tmp_path) -> None: + db = tmp_path / "t.sqlite" + store = SqliteStore(str(db)) + t = store.create_tenant("T") + tid = str(t["id"]) + store.create_user_account( + tenant_id=tid, + username="u", + display_name="U", + role="owner", + password_hash="x", + is_active=True, + ) + uid = str(store.get_user_by_username(tenant_id=tid, username="u")["id"]) + s = store.create_session_for_user(title="s", tenant_id=tid, user_id=uid) + sid = str(s.id) + tu = uuid.uuid4().hex + store.add_message(session_id=sid, role="user", content="hi", turn_uuid=tu, event_type="user_text") + assert persist_assistant_text_if_turn_missing( + store=store, + session_id=sid, + turn_uuid=tu, + final_text="fallback body", + log_prefix="test_fallback", + ) + msgs = store.get_messages(sid, limit=20) + assert any(getattr(m, "role", "") == "assistant" and "fallback" in str(getattr(m, "content", "")) for m in msgs) + + +def test_fallback_skips_when_present(tmp_path) -> None: + db = tmp_path / "u.sqlite" + store = SqliteStore(str(db)) + t = store.create_tenant("T2") + tid = str(t["id"]) + store.create_user_account( + tenant_id=tid, + username="u2", + display_name="U", + role="owner", + password_hash="x", + is_active=True, + ) + uid = str(store.get_user_by_username(tenant_id=tid, username="u2")["id"]) + s = store.create_session_for_user(title="s", tenant_id=tid, user_id=uid) + sid = str(s.id) + tu = uuid.uuid4().hex + store.add_message(session_id=sid, role="user", content="hi", turn_uuid=tu, event_type="user_text") + store.add_message(session_id=sid, role="assistant", content="already", turn_uuid=tu, event_type="assistant_text") + assert not persist_assistant_text_if_turn_missing( + store=store, + session_id=sid, + turn_uuid=tu, + final_text="would duplicate", + log_prefix="test_skip", + ) diff --git a/tests/test_pg_compat.py b/tests/test_pg_compat.py new file mode 100644 index 00000000..fd6cfb10 --- /dev/null +++ b/tests/test_pg_compat.py @@ -0,0 +1,46 @@ +"""Guardrails for SQLite→PostgreSQL SQL rewriting used by SqliteStore.""" + +from __future__ import annotations + +from svc.persistence import pg_compat +from svc.persistence.pg_adapter import normalize_psycopg_conninfo + + +def test_normalize_psycopg_conninfo_strips_sqlalchemy_driver() -> None: + u = "postgresql+psycopg://user:pass@127.0.0.1:5432/oclaw" + assert normalize_psycopg_conninfo(u) == "postgresql://user:pass@127.0.0.1:5432/oclaw" + + +def test_insert_or_replace_ui_session_owner_rewritten() -> None: + sql = """INSERT OR REPLACE INTO ui_session_owner(session_id, tenant_id, user_id, created_at) + VALUES (?, ?, ?, ?)""" + adapted = pg_compat.adapt_sql_for_postgres(sql) + assert "INSERT OR REPLACE" not in adapted.upper() + assert "ON CONFLICT" in adapted.upper() + + +def test_llm_profile_insert_or_ignore_rewritten() -> None: + sql = """INSERT OR IGNORE INTO llm_profile + (id, name, mode, model, base_url, api_key, updated_at, is_builtin, hide_in_ui, owner_user_id) + VALUES (?, ?, 'ollama', ?, ?, NULL, ?, 1, 0, NULL)""" + adapted = pg_compat.adapt_sql_for_postgres(sql) + assert "INSERT OR IGNORE" not in adapted.upper() + assert "ON CONFLICT (id) DO NOTHING" in adapted + + +def test_role_permission_insert_or_ignore_rewritten() -> None: + sql = """INSERT OR IGNORE INTO role_permission(role, permission, created_at) + VALUES (?, ?, ?)""" + adapted = pg_compat.adapt_sql_for_postgres(sql) + assert "INSERT OR IGNORE" not in adapted.upper() + assert "ON CONFLICT (role, permission) DO NOTHING" in adapted + + +def test_scrub_nul_bytes_from_text() -> None: + assert pg_compat.scrub_nul_bytes_from_text(None) is None + assert pg_compat.scrub_nul_bytes_from_text("ok") == "ok" + assert pg_compat.scrub_nul_bytes_from_text("a\x00b") == "ab" + + +def test_scrub_nul_bytes_from_jsonable_nested() -> None: + assert pg_compat.scrub_nul_bytes_from_jsonable({"x": "y\x00z"}) == {"x": "yz"} diff --git a/tests/test_sa_admin_user_stats_tool_health.py b/tests/test_sa_admin_user_stats_tool_health.py new file mode 100644 index 00000000..212f5e9e --- /dev/null +++ b/tests/test_sa_admin_user_stats_tool_health.py @@ -0,0 +1,78 @@ +"""SA migration: list_admin_user_stats + list_session_tool_health.""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_admin_th.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_list_admin_user_stats_tokens_sessions_logins(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("StatsT") + u = s.create_user(tenant_id=t["id"], display_name="Stats User", role="member") + sess = s.create_session_for_user(title="Sess", tenant_id=t["id"], user_id=u["id"]) + s.add_message(sess.id, "user", "hi") + s.add_trace_event( + session_id=sess.id, + trace_id="tr1", + span_id="sp1", + parent_span_id=None, + event_type="llm", + payload={"prompt_tokens_est": 12, "response_tokens_est": 3}, + ) + exp = (datetime.now(timezone.utc) + timedelta(days=1)).isoformat() + s.create_auth_session( + session_token_hash="tok_stats_" + sess.id[:8], + tenant_id=t["id"], + user_id=u["id"], + role="member", + expires_at=exp, + ) + total, users, totals = s.list_admin_user_stats(tenant_id=t["id"], limit=50, offset=0) + assert total >= 1 + row = next(x for x in users if x["user_id"] == u["id"]) + assert row["sessions_count"] >= 1 + assert row["total_tokens_est"] == 15 + assert totals["users_count"] == total + assert totals["active_sessions_30m"] >= 1 + assert totals["active_logins_30m"] >= 1 + + +def test_sa_list_session_tool_health_warn_then_ok(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("TH") + s.add_message(sess.id, "assistant", "only assistant") + rows = s.list_session_tool_health(session_id=sess.id, limit=10) + assert len(rows) == 1 + assert rows[0]["status"] == "warn_no_tool_calls" + assert rows[0]["assistant_count"] == 1 + s.add_tool_log(sess.id, "grep", {}, {"ok": True}) + rows2 = s.list_session_tool_health(session_id=sess.id, limit=10) + assert rows2[0]["status"] == "ok" + assert rows2[0]["tool_count"] >= 1 + + +def test_sa_list_session_tool_health_mcp_count(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("MCP") + s.add_tool_log(sess.id, "mcp__srv1__ping", {}, {"ok": True}) + rows = s.list_session_tool_health(session_id=sess.id, limit=10) + assert rows[0]["mcp_tool_count"] == 1 diff --git a/tests/test_sa_app_settings.py b/tests/test_sa_app_settings.py new file mode 100644 index 00000000..d7e58ef8 --- /dev/null +++ b/tests/test_sa_app_settings.py @@ -0,0 +1,41 @@ +"""SQLAlchemy slice tests: app_setting repository (phase-1 migration).""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_slice.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_app_setting_roundtrip(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + s.set_setting("sa_k", "sa_v") + assert s.get_setting("sa_k") == "sa_v" + s.set_setting("sa_k", "sa_v2") + assert s.get_setting("sa_k") == "sa_v2" + s.delete_setting("sa_k") + assert s.get_setting("sa_k") is None + + +def test_sa_app_setting_get_ignores_secret(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + s.set_secret("sec_k", "plain") + assert s.get_setting("sec_k") is None + assert s.get_secret("sec_k") == "plain" + s.delete_setting("sec_k") + assert s.get_secret("sec_k") is None diff --git a/tests/test_sa_app_users.py b/tests/test_sa_app_users.py new file mode 100644 index 00000000..f2de2675 --- /dev/null +++ b/tests/test_sa_app_users.py @@ -0,0 +1,132 @@ +"""SA migration: app_user create + lookups.""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_app_users.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_create_user_username_suffix(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("UApp") + u1 = s.create_user(tenant_id=t["id"], display_name="Same Name", role="member") + u2 = s.create_user(tenant_id=t["id"], display_name="Same Name", role="member") + assert u1["username"] != u2["username"] + assert u2["username"].startswith(u1["username"]) + + +def test_sa_get_user_by_id_and_username(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("GApp") + u = s.create_user(tenant_id=t["id"], display_name="Lookup", role="owner") + by_id = s.get_user_by_id(tenant_id=t["id"], user_id=u["id"]) + assert by_id is not None + assert by_id["id"] == u["id"] + by_un = s.get_user_by_username(tenant_id=t["id"], username=u["username"]) + assert by_un is not None + assert by_un["id"] == u["id"] + + +def test_sa_create_user_account_and_get_global(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("AccT") + row = s.create_user_account( + tenant_id=t["id"], + username="acctester", + display_name="Acc", + role="member", + password_hash="h", + is_active=True, + ) + assert row.get("username") == "acctester" + g = s.get_user_by_username_global(username="acctester") + assert g is not None + assert g["tenant_id"] == t["id"] + + +def test_sa_list_users_filter_order_and_flags(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("LU") + u1 = s.create_user(tenant_id=t["id"], display_name="Alpha User", role="member") + u2 = s.create_user(tenant_id=t["id"], display_name="Beta", role="owner") + rows = s.list_users(tenant_id=t["id"], limit=10, offset=0) + assert {u1["id"], u2["id"]}.issubset({r["id"] for r in rows}) + assert rows[0]["id"] == u2["id"] + + filtered = s.list_users(tenant_id=t["id"], q="alpha") + assert len(filtered) == 1 + assert filtered[0]["id"] == u1["id"] + + +def test_sa_list_users_include_inactive(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("IA") + s.create_user(tenant_id=t["id"], display_name="Active", role="member") + inac = s.create_user_account( + tenant_id=t["id"], + username="inac1", + display_name="Inactive", + role="member", + password_hash="x", + is_active=False, + ) + assert any(r["id"] == inac["id"] for r in s.list_users(tenant_id=t["id"], include_inactive=True)) + assert not any(r["id"] == inac["id"] for r in s.list_users(tenant_id=t["id"], include_inactive=False)) + + +def test_sa_list_users_wecom_ids_and_update_delete(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("WC") + u = s.create_user(tenant_id=t["id"], display_name="We", role="member") + ts = "2026-01-01T00:00:00Z" + with s._connect() as conn: + conn.execute( + """ + INSERT INTO channel_identity (tenant_id, channel, external_user_id, user_id, created_at) + VALUES (?, 'wecom', 'ext-a', ?, ?) + """, + (t["id"], u["id"], ts), + ) + conn.execute( + """ + INSERT INTO channel_identity (tenant_id, channel, external_user_id, user_id, created_at) + VALUES (?, 'wecom', 'ext-b', ?, ?) + """, + (t["id"], u["id"], ts), + ) + row = next(r for r in s.list_users(tenant_id=t["id"]) if r["id"] == u["id"]) + assert row["wecom_linked"] is True + assert row["channel_linked"] is True + eids = row["wecom_external_user_ids"] + assert "ext-a" in eids and "ext-b" in eids + + assert s.update_user_account(tenant_id=t["id"], user_id=u["id"], display_name="We2", role="owner") is True + ref = s.get_user_by_id(tenant_id=t["id"], user_id=u["id"]) + assert ref is not None + assert ref["display_name"] == "We2" and ref["role"] == "owner" + + assert s.delete_user_account(tenant_id=t["id"], user_id=u["id"]) == 1 + assert s.get_user_by_id(tenant_id=t["id"], user_id=u["id"]) is None + + +def test_sa_update_user_account_noop(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("NP") + u = s.create_user(tenant_id=t["id"], display_name="N", role="member") + assert s.update_user_account(tenant_id=t["id"], user_id=u["id"]) is False diff --git a/tests/test_sa_auth_sessions.py b/tests/test_sa_auth_sessions.py new file mode 100644 index 00000000..46f13e28 --- /dev/null +++ b/tests/test_sa_auth_sessions.py @@ -0,0 +1,75 @@ +"""SQLAlchemy slice tests: auth_session repository (phase-2 migration).""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +import uuid + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_auth.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def _future_expires() -> str: + return (datetime.now(timezone.utc) + timedelta(days=1)).isoformat() + + +def test_sa_auth_session_roundtrip(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("AuthSlice") + u = s.create_user(tenant_id=t["id"], display_name="Member", role="member") + h = "tokhash_" + uuid.uuid4().hex + exp = _future_expires() + s.create_auth_session( + session_token_hash=h, + tenant_id=t["id"], + user_id=u["id"], + role="member", + expires_at=exp, + ) + row = s.get_auth_session(session_token_hash=h) + assert row is not None + assert row["tenant_id"] == t["id"] + assert row["user_id"] == u["id"] + assert not row.get("revoked_at") + s.touch_auth_session(session_token_hash=h) + row2 = s.get_auth_session(session_token_hash=h) + assert row2 is not None + assert row2["last_seen_at"] is not None + assert s.revoke_auth_session(session_token_hash=h) == 1 + assert s.revoke_auth_session(session_token_hash=h) == 0 + row3 = s.get_auth_session(session_token_hash=h) + assert row3 is not None + assert row3.get("revoked_at") + + +def test_sa_auth_session_revoke_all(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("AuthSlice2") + u = s.create_user(tenant_id=t["id"], display_name="M2", role="member") + exp = _future_expires() + for i in range(2): + s.create_auth_session( + session_token_hash=f"h{i}_" + uuid.uuid4().hex, + tenant_id=t["id"], + user_id=u["id"], + role="member", + expires_at=exp, + ) + n = s.revoke_all_auth_sessions() + assert n == 2 diff --git a/tests/test_sa_chat_messages.py b/tests/test_sa_chat_messages.py new file mode 100644 index 00000000..4d4b10ea --- /dev/null +++ b/tests/test_sa_chat_messages.py @@ -0,0 +1,91 @@ +"""SQLAlchemy slice tests: chat_message hot paths (phase-4 migration).""" + +from __future__ import annotations + +import json + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_msg.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_chat_message_add_list_meta(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("M") + m1 = s.add_message(sess.id, "user", "hello") + assert m1.id > 0 + m2 = s.add_message(sess.id, "assistant", "hi") + assert m2.id > m1.id + meta = s.get_session_messages_meta(sess.id) + assert meta.session_id == sess.id + assert meta.message_count == 2 + assert meta.last_message_id == m2.id + assert s.count_messages(sess.id) == 2 + assert s.get_last_message_id(sess.id) == m2.id + rows = s.get_messages(sess.id, limit=10) + assert [x.role for x in rows] == ["user", "assistant"] + assert rows[0].content == "hello" + after = s.get_messages_after_id(session_id=sess.id, after_id=m1.id, limit=10) + assert len(after) == 1 + assert after[0].id == m2.id + + +def test_sa_chat_message_delete_refreshes_session_ts(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("D") + a = s.add_message(sess.id, "user", "a", timestamp="2020-01-01T00:00:01+00:00") + b = s.add_message(sess.id, "user", "b", timestamp="2020-01-01T00:00:02+00:00") + assert s.get_session(sess.id).last_message_at is not None + assert s.delete_message(session_id=sess.id, message_id=b.id) is True + meta = s.get_session_messages_meta(sess.id) + assert meta.message_count == 1 + assert meta.last_message_id == a.id + assert s.delete_message(session_id=sess.id, message_id=999999) is False + + +def test_sa_chat_message_update_content(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("U") + m = s.add_message(sess.id, "assistant", "old", event_payload={"k": 1}) + ok = s.update_message_content(session_id=sess.id, message_id=m.id, content="new", event_payload={"k": 2}) + assert ok is True + row = s.get_messages(sess.id, limit=5)[0] + assert row.content == "new" + assert json.loads(row.event_payload or "{}")["k"] == 2 + + +def test_sa_chat_message_scrubs_nul_byte(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("N") + m = s.add_message(sess.id, "user", "a\x00b", event_payload={"x": "y\x00z"}) + rows = s.get_messages(sess.id, limit=5) + assert rows[0].content == "ab" + ep = json.loads(rows[0].event_payload or "{}") + assert ep["x"] == "yz" + + +def test_sa_chat_message_tool_window_prepend(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("T") + asst = s.add_message(sess.id, "assistant", "call", tool_calls='{"x":1}') + tool_calls = json.dumps({"assistant_message_id": asst.id}, ensure_ascii=False) + s.add_message(sess.id, "tool", "result", tool_calls=tool_calls) + win = s.get_messages(sess.id, limit=1) + assert len(win) == 2 + assert win[0].id == asst.id + assert win[1].role == "tool" diff --git a/tests/test_sa_chat_sessions.py b/tests/test_sa_chat_sessions.py new file mode 100644 index 00000000..f2300010 --- /dev/null +++ b/tests/test_sa_chat_sessions.py @@ -0,0 +1,52 @@ +"""SQLAlchemy slice tests: chat_session list + CRUD (phase-3 migration).""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_chat.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_chat_session_crud(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + meta0 = s.get_sessions_list_meta() + assert meta0.session_count == 0 + sess = s.create_session("Hi") + assert s.get_session(sess.id) is not None + assert s.get_session(sess.id).title == "Hi" + assert s.count_sessions() == 1 + s.rename_session(sess.id, "Hi2") + assert s.get_session(sess.id).title == "Hi2" + listed = s.list_sessions(limit=5, offset=0) + assert len(listed) == 1 + s.delete_session(sess.id) + assert s.get_session(sess.id) is None + + +def test_sa_chat_session_user_scoped(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("ChatT") + u = s.create_user(tenant_id=t["id"], display_name="ChatU", role="member") + sess = s.create_session_for_user(title="Owned", tenant_id=t["id"], user_id=u["id"]) + rows = s.list_sessions_for_user(tenant_id=t["id"], user_id=u["id"], limit=10, offset=0) + assert len(rows) == 1 + assert rows[0].id == sess.id + assert s.get_session_for_user(session_id=sess.id, tenant_id=t["id"], user_id=u["id"]) is not None + assert s.get_session_in_tenant(session_id=sess.id, tenant_id=t["id"]) is not None + assert s.delete_session_for_user(session_id=sess.id, tenant_id=t["id"], user_id=u["id"]) is True + assert s.get_session(sess.id) is None diff --git a/tests/test_sa_fork_trim_admin_sessions.py b/tests/test_sa_fork_trim_admin_sessions.py new file mode 100644 index 00000000..00fa1738 --- /dev/null +++ b/tests/test_sa_fork_trim_admin_sessions.py @@ -0,0 +1,85 @@ +"""SA migration: fork_session, trim_messages, list_admin_sessions.""" + +from __future__ import annotations + +import json + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_fork.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_fork_session_remaps_tool_assistant_id(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + src = s.create_session("Src") + asst = s.add_message(src.id, "assistant", "a", tool_calls="{}") + tc = json.dumps({"assistant_message_id": asst.id}, ensure_ascii=False) + s.add_message(src.id, "tool", "t", tool_calls=tc) + forked = s.fork_session(src.id, up_to_message_id=int(asst.id) + 1, title="Forked") + assert forked.id != src.id + msgs = s.get_messages(forked.id, limit=10) + assert len(msgs) == 2 + assert msgs[0].role == "assistant" + assert msgs[1].role == "tool" + meta = json.loads(msgs[1].tool_calls or "{}") + assert int(meta.get("assistant_message_id") or 0) == msgs[0].id + + +def test_sa_fork_session_bad_anchor(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + src = s.create_session("S2") + m = s.add_message(src.id, "user", "x") + n_before = s.count_sessions() + with pytest.raises(ValueError, match="message not in session"): + s.fork_session(src.id, up_to_message_id=int(m.id) + 99, title="Bad") + assert s.count_sessions() == n_before + + +def test_sa_trim_messages_keep_last(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("Trim") + for i in range(5): + s.add_message(sess.id, "user", f"m{i}") + s.trim_messages(sess.id, keep_last=2) + assert s.count_messages(sess.id) == 2 + rows = s.get_messages(sess.id, limit=10) + assert [x.content for x in rows] == ["m3", "m4"] + + +def test_sa_trim_messages_zero_deletes_session(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("Z") + s.add_message(sess.id, "user", "u") + s.trim_messages(sess.id, keep_last=0) + assert s.get_session(sess.id) is None + + +def test_sa_list_admin_sessions(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("AdmT") + u = s.create_user(tenant_id=t["id"], display_name="Alice Admin", role="member") + sess = s.create_session_for_user(title="S1", tenant_id=t["id"], user_id=u["id"]) + s.add_message(sess.id, "user", "hello") + total, rows = s.list_admin_sessions(tenant_id=t["id"], limit=10, offset=0) + assert total >= 1 + hit = next(x for x in rows if x["session_id"] == sess.id) + assert hit["message_count"] == 1 + assert "alice" in hit["display_name"].lower() or "alice" in hit["username"].lower() + t2, rows2 = s.list_admin_sessions(tenant_id=t["id"], q="S1", limit=10, offset=0) + assert t2 >= 1 + assert any(x["session_id"] == sess.id for x in rows2) diff --git a/tests/test_sa_prune_orphans.py b/tests/test_sa_prune_orphans.py new file mode 100644 index 00000000..d1cad0e5 --- /dev/null +++ b/tests/test_sa_prune_orphans.py @@ -0,0 +1,54 @@ +"""SA housekeeping: delete orphan chat_message / tool_log rows.""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_prune.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_delete_orphan_chat_message_and_tool_log(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + with s._connect() as conn: + conn.execute("PRAGMA foreign_keys = OFF") + conn.execute( + "INSERT INTO chat_message (session_id, role, content, timestamp) VALUES (?, ?, ?, ?)", + ("ghost-sid", "user", "orphan", "2020-01-01T00:00:00+00:00"), + ) + conn.execute( + """ + INSERT INTO tool_log (session_id, tool_name, specialist, args, result, timestamp, duration_ms) + VALUES (?, ?, '', '{}', '{}', ?, NULL) + """, + ("ghost-sid", "t", "2020-01-01T00:00:01+00:00"), + ) + conn.execute("PRAGMA foreign_keys = ON") + n_msg = s._chat_messages_repo().delete_messages_where_session_missing() + n_tl = s._tool_log_queries_repo().delete_tool_logs_where_session_missing() + assert n_msg >= 1 + assert n_tl >= 1 + + +def test_sa_prune_runs_with_sa_deletes(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + s.create_session("keep") + with s._connect() as conn: + s._prune_rows_for_missing_chat_session(conn) + # second call should not raise + with s._connect() as conn: + s._prune_rows_for_missing_chat_session(conn) diff --git a/tests/test_sa_tenant_bind_code.py b/tests/test_sa_tenant_bind_code.py new file mode 100644 index 00000000..87604235 --- /dev/null +++ b/tests/test_sa_tenant_bind_code.py @@ -0,0 +1,48 @@ +"""SA migration: tenant + bind_code.""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_tenant_bc.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_tenant_crud(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("Acme") + listed = s.list_tenants(limit=50) + assert any(x["id"] == t["id"] for x in listed) + assert s.delete_tenant(tenant_id=t["id"]) == 1 + assert s.delete_tenant(tenant_id=t["id"]) == 0 + + +def test_sa_bind_code_list_and_consume(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("BindCo") + s.create_bind_code(tenant_id=t["id"], role="member", code="CODE12345") + rows = s.list_bind_codes(tenant_id=t["id"], limit=10) + assert len(rows) == 1 + assert rows[0]["code"] == "CODE12345" + assert rows[0]["used_at"] is None + out = s.consume_bind_code(code="CODE12345", channel="wecom", external_user_id="wx-1", display_name="Ext") + assert out is not None + assert out["tenant_id"] == t["id"] + rows2 = s.list_bind_codes(tenant_id=t["id"], limit=10) + assert rows2[0]["used_at"] is not None + assert rows2[0]["used_by_external_user_id"] == "wx-1" + assert s.consume_bind_code(code="CODE12345", channel="wecom", external_user_id="wx-2") is None diff --git a/tests/test_sa_tool_log_trace.py b/tests/test_sa_tool_log_trace.py new file mode 100644 index 00000000..d94d869d --- /dev/null +++ b/tests/test_sa_tool_log_trace.py @@ -0,0 +1,121 @@ +"""SA migration: tool_log MCP queries + trace_event writes/lists.""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_tl_tr.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_mcp_tool_summaries_and_call_logs(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("MCP2") + s.add_tool_log(sess.id, "mcp__srv99__x", {}, {"ok": True}, specialist="sp1") + s.add_tool_log(sess.id, "mcp__srv99__x", {}, {"ok": True}, specialist="sp1") + agg = s.list_mcp_tool_aggregate_usage() + assert "mcp__srv99__x" in agg + assert agg["mcp__srv99__x"]["count"] == 2 + summ = s.list_mcp_tool_usage_summary(limit=50) + hit = next(x for x in summ if x["tool_name"] == "mcp__srv99__x") + assert hit["count"] == 2 + assert hit["server_id"] == "srv99" + logs = s.list_mcp_tool_call_logs(server_id="srv99", limit=10) + assert len(logs) == 2 + assert logs[0]["session_id"] == sess.id + + +def test_sa_trace_roundtrip_batch_and_window(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("TR") + s.add_trace_event( + session_id=sess.id, + trace_id="t1", + span_id="s0", + parent_span_id=None, + event_type="noise", + payload={"a": 1}, + ) + s.add_trace_events_batch( + [ + { + "session_id": sess.id, + "trace_id": "t1", + "span_id": "s1", + "parent_span_id": None, + "event_type": "turn_started", + "payload": {}, + }, + { + "session_id": sess.id, + "trace_id": "t1", + "span_id": "s2", + "parent_span_id": None, + "event_type": "turn_finished", + "payload": {}, + }, + ] + ) + desc = s.list_trace_events(session_id=sess.id, limit=10) + assert len(desc) == 3 + asc_rows = s.list_trace_events_for_trace(session_id=sess.id, trace_id="t1", limit=50) + assert len(asc_rows) == 3 + assert asc_rows[0]["event_type"] == "noise" + st, en = s.get_turn_time_window(session_id=sess.id, trace_id="t1") + assert st is not None and en is not None + + +def test_sa_add_tool_log_get_tool_logs(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("TL") + s.add_tool_log(sess.id, "my_tool", {"a": 1}, {"ok": True}, specialist="spec", duration_ms=42) + logs = s.get_tool_logs(sess.id, limit=10) + assert len(logs) == 1 + assert logs[0]["tool_name"] == "my_tool" + assert logs[0]["specialist"] == "spec" + assert logs[0]["args"] == {"a": 1} + assert logs[0]["result"] == {"ok": True} + assert logs[0]["duration_ms"] == 42 + + +def test_sa_list_messages_in_time_window(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + sess = s.create_session("Win") + t0 = "2024-01-01T10:00:00+00:00" + t1 = "2024-01-01T10:05:00+00:00" + t2 = "2024-01-01T10:10:00+00:00" + s.add_message(sess.id, "user", "a", timestamp=t0) + s.add_message(sess.id, "user", "b", timestamp=t1) + s.add_message(sess.id, "user", "c", timestamp=t2) + rows = s.list_messages_in_time_window(session_id=sess.id, start_ts=t0, end_ts=t1, limit=50) + assert len(rows) == 2 + assert [x["content"] for x in rows] == ["a", "b"] + + +def test_sa_move_tool_logs_to_session(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + src = s.create_session("SrcS") + dst = s.create_session("DstS") + s.add_tool_log(src.id, "t_move", {}, {"ok": 1}) + n = s.move_tool_logs_to_session(from_session_id=src.id, to_session_id=dst.id) + assert n == 1 + assert s.get_tool_logs(src.id, limit=10) == [] + logs = s.get_tool_logs(dst.id, limit=10) + assert len(logs) == 1 + assert logs[0]["tool_name"] == "t_move" + assert s.move_tool_logs_to_session(from_session_id=src.id, to_session_id=dst.id) == 0 + assert s.move_tool_logs_to_session(from_session_id="", to_session_id=dst.id) == 0 diff --git a/tests/test_sa_ui_session_owner.py b/tests/test_sa_ui_session_owner.py new file mode 100644 index 00000000..63a322cb --- /dev/null +++ b/tests/test_sa_ui_session_owner.py @@ -0,0 +1,96 @@ +"""SA migration: ui_session_owner upsert + backfill.""" + +from __future__ import annotations + +import pytest + +from svc.persistence.db.engine import clear_assistant_engine_cache +from svc.persistence.sqlite_store import SqliteStore + + +@pytest.fixture +def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path): + monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False) + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + dbfile = tmp_path / "sa_ui_owner.sqlite" + monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile)) + clear_assistant_engine_cache() + s = SqliteStore(str(dbfile)) + try: + yield s + finally: + clear_assistant_engine_cache() + + +def test_sa_create_session_for_user_and_get_owner(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("UIO") + u = s.create_user(tenant_id=t["id"], display_name="U", role="member") + sess = s.create_session_for_user(title="T", tenant_id=t["id"], user_id=u["id"]) + row = s.get_ui_session_owner(session_id=sess.id) + assert row is not None + assert row["tenant_id"] == t["id"] + assert row["user_id"] == u["id"] + s.ensure_ui_session_owner(session_id=sess.id, tenant_id=t["id"], user_id=u["id"]) + row2 = s.get_ui_session_owner(session_id=sess.id) + assert row2 == row + + +def test_sa_upsert_replace_changes_owner(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("UIO2") + u1 = s.create_user(tenant_id=t["id"], display_name="A", role="member") + u2 = s.create_user(tenant_id=t["id"], display_name="B", role="member") + sess = s.create_session_for_user(title="X", tenant_id=t["id"], user_id=u1["id"]) + s._ui_session_owner_repo().upsert_replace( + session_id=sess.id, + tenant_id=t["id"], + user_id=u2["id"], + created_at="2099-01-01T00:00:00+00:00", + ) + row = s.get_ui_session_owner(session_id=sess.id) + assert row is not None + assert row["user_id"] == u2["id"] + + +def test_sa_backfill_orphan_sessions(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("BF") + u = s.create_user(tenant_id=t["id"], display_name="BFU", role="member") + orphan = s.create_session("orphan title") + assert s.get_ui_session_owner(session_id=orphan.id) is None + n = s.backfill_orphan_chat_sessions_for_user(tenant_id=t["id"], user_id=u["id"]) + assert n >= 1 + row = s.get_ui_session_owner(session_id=orphan.id) + assert row is not None + assert row["user_id"] == u["id"] + + +def test_sa_backfill_ui_session_owner_from_channel_v2(fresh_sqlite_store: SqliteStore) -> None: + s = fresh_sqlite_store + t = s.create_tenant("CH") + u = s.create_user(tenant_id=t["id"], display_name="CHU", role="member") + sess = s.create_session("ch sess") + ts = "2025-01-01T00:00:00+00:00" + with s._connect() as conn: + conn.execute( + """ + INSERT INTO channel_identity_v2 + (tenant_id, channel, account_id, external_user_id, user_id, created_at) + VALUES (?, 'wecom', 'acc1', 'extu', ?, ?) + """, + (t["id"], u["id"], ts), + ) + conn.execute( + """ + INSERT INTO channel_session_v2 + (tenant_id, channel, account_id, external_chat_id, external_user_id, session_id, created_at) + VALUES (?, 'wecom', 'acc1', 'chat1', 'extu', ?, ?) + """, + (t["id"], sess.id, ts), + ) + n = s.backfill_ui_session_owner_from_channel_v2() + assert n >= 1 + row = s.get_ui_session_owner(session_id=sess.id) + assert row is not None + assert row["user_id"] == u["id"]