mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 00:40:45 +08:00
feat(persistence): PostgreSQL assistant store, chat persist fixes, gateway scripts
- Add SQLAlchemy Core repos, pg adapter/compat, assistant_store factory, Alembic bootstrap and migration/cutover scripts. - Harden chat_message writes (NUL scrub for PG), turn_uuid on attempt failure, WS turn_runner fallbacks and gateway executed_turn_uuid init. - start_gateway: log paths, PS7 stderr handling via cmd, background stdout/stderr redirect; runtime assistant_runtime_log_dir export. - Ops: clear_all_chat_sessions with PG-only --postgresql and env-gated wipe; clear_postgres_chat_sessions.ps1. - Tests: SA repos, pg compat, persist fallback, smoke env isolation; CI and docs touch-ups. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
2b32d11f43
commit
d14e9d3596
103 changed files with 7574 additions and 1641 deletions
46
.github/workflows/ci.yml
vendored
46
.github/workflows/ci.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 <postgresql+psycopg或 postgresql 连接串>]
|
||||
# 默认要求 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=
|
||||
|
|
|
|||
37
alembic.ini
Normal file
37
alembic.ini
Normal file
|
|
@ -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
|
||||
67
assistant_migrations/env.py
Normal file
67
assistant_migrations/env.py
Normal file
|
|
@ -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()
|
||||
26
assistant_migrations/script.py.mako
Normal file
26
assistant_migrations/script.py.mako
Normal file
|
|
@ -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"}
|
||||
34
assistant_migrations/versions/001_assistant_pg_initial.py
Normal file
34
assistant_migrations/versions/001_assistant_pg_initial.py
Normal file
|
|
@ -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.")
|
||||
|
|
@ -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 路径
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 ? `<div style="margin-top:8px;"><span class="muted">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".
|
||||
} 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 (_) {}
|
||||
await loadMessagesForActive();
|
||||
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.
|
||||
}
|
||||
if (needWsFallback) {
|
||||
if (n < 0) {
|
||||
try {
|
||||
if (streamRow && streamRow.parentNode) streamRow.remove();
|
||||
} catch (_) {}
|
||||
await loadMessagesForActive();
|
||||
scrollMessagesToBottom(true);
|
||||
}
|
||||
ok = await appendFinalAssistant(payload.message, fbLine);
|
||||
} else if (n > 0) {
|
||||
ok = true;
|
||||
} else {
|
||||
ok = n === 0;
|
||||
}
|
||||
scrollMessagesToBottom(true);
|
||||
} else {
|
||||
ok = await appendFinalAssistant(payload.message, chatStream || extractWsAssistantText(payload.message || {}));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"] = {}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
46
runtime/chat/persist_terminal_fallback.py
Normal file
46
runtime/chat/persist_terminal_fallback.py
Normal file
|
|
@ -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"]
|
||||
|
|
@ -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 ""),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 ``<parent of sqlite db_path>)/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():
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
"$@"
|
||||
147
runtime/operations/scripts/clear_all_chat_sessions.py
Normal file
147
runtime/operations/scripts/clear_all_chat_sessions.py
Normal file
|
|
@ -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())
|
||||
71
runtime/operations/scripts/clear_postgres_chat_sessions.ps1
Normal file
71
runtime/operations/scripts/clear_postgres_chat_sessions.ps1
Normal file
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
128
runtime/operations/scripts/cutover_sqlite_to_postgresql.ps1
Normal file
128
runtime/operations/scripts/cutover_sqlite_to_postgresql.ps1
Normal file
|
|
@ -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
|
||||
67
runtime/operations/scripts/cutover_sqlite_to_postgresql.sh
Normal file
67
runtime/operations/scripts/cutover_sqlite_to_postgresql.sh
Normal file
|
|
@ -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: <repo>/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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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] = [
|
||||
|
|
|
|||
251
runtime/operations/scripts/live_chat_probe.py
Normal file
251
runtime/operations/scripts/live_chat_probe.py
Normal file
|
|
@ -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 <id> --token <jwt>
|
||||
|
||||
可选:加 ``--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())
|
||||
|
|
@ -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()
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 <bot_id> <bot_secret>")
|
||||
return 2
|
||||
|
||||
store = SqliteStore(db_path())
|
||||
store = get_assistant_store()
|
||||
if len(args) < 2:
|
||||
print("error=missing_required_args")
|
||||
return 2
|
||||
|
|
|
|||
178
runtime/operations/scripts/smoke_admin_chat_postgres.py
Normal file
178
runtime/operations/scripts/smoke_admin_chat_postgres.py
Normal file
|
|
@ -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())
|
||||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()}")
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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": "",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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())]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
4
scripts/clear_postgres_chat_sessions.ps1
Normal file
4
scripts/clear_postgres_chat_sessions.ps1
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
57
svc/config/database.py
Normal file
57
svc/config/database.py
Normal file
|
|
@ -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"]
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
57
svc/persistence/assistant_store.py
Normal file
57
svc/persistence/assistant_store.py
Normal file
|
|
@ -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"]
|
||||
192
svc/persistence/assistant_store_protocol.py
Normal file
192
svc/persistence/assistant_store_protocol.py
Normal file
|
|
@ -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"]
|
||||
43
svc/persistence/db/__init__.py
Normal file
43
svc/persistence/db/__init__.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
72
svc/persistence/db/engine.py
Normal file
72
svc/persistence/db/engine.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
171
svc/persistence/db/tables.py
Normal file
171
svc/persistence/db/tables.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
486
svc/persistence/ddl/postgresql_bootstrap.sql
Normal file
486
svc/persistence/ddl/postgresql_bootstrap.sql
Normal file
|
|
@ -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);
|
||||
68
svc/persistence/pg_adapter.py
Normal file
68
svc/persistence/pg_adapter.py
Normal file
|
|
@ -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"]
|
||||
230
svc/persistence/pg_compat.py
Normal file
230
svc/persistence/pg_compat.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
30
svc/persistence/sa_repos/__init__.py
Normal file
30
svc/persistence/sa_repos/__init__.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
167
svc/persistence/sa_repos/admin_user_stats.py
Normal file
167
svc/persistence/sa_repos/admin_user_stats.py
Normal file
|
|
@ -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"]
|
||||
141
svc/persistence/sa_repos/app_settings.py
Normal file
141
svc/persistence/sa_repos/app_settings.py
Normal file
|
|
@ -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"]
|
||||
296
svc/persistence/sa_repos/app_users.py
Normal file
296
svc/persistence/sa_repos/app_users.py
Normal file
|
|
@ -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"]
|
||||
102
svc/persistence/sa_repos/auth_sessions.py
Normal file
102
svc/persistence/sa_repos/auth_sessions.py
Normal file
|
|
@ -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"]
|
||||
481
svc/persistence/sa_repos/chat_messages.py
Normal file
481
svc/persistence/sa_repos/chat_messages.py
Normal file
|
|
@ -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"]
|
||||
424
svc/persistence/sa_repos/chat_sessions.py
Normal file
424
svc/persistence/sa_repos/chat_sessions.py
Normal file
|
|
@ -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"]
|
||||
77
svc/persistence/sa_repos/session_tool_health.py
Normal file
77
svc/persistence/sa_repos/session_tool_health.py
Normal file
|
|
@ -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"]
|
||||
116
svc/persistence/sa_repos/tenant_bind_code.py
Normal file
116
svc/persistence/sa_repos/tenant_bind_code.py
Normal file
|
|
@ -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"]
|
||||
133
svc/persistence/sa_repos/tool_log_queries.py
Normal file
133
svc/persistence/sa_repos/tool_log_queries.py
Normal file
|
|
@ -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"]
|
||||
108
svc/persistence/sa_repos/trace_events.py
Normal file
108
svc/persistence/sa_repos/trace_events.py
Normal file
|
|
@ -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"]
|
||||
193
svc/persistence/sa_repos/ui_session_owner.py
Normal file
193
svc/persistence/sa_repos/ui_session_owner.py
Normal file
|
|
@ -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"]
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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", "")) == "几点了"
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
84
tests/test_database_backend.py
Normal file
84
tests/test_database_backend.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
86
tests/test_migrate_assistant_sqlite_to_pg.py
Normal file
86
tests/test_migrate_assistant_sqlite_to_pg.py
Normal file
|
|
@ -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"
|
||||
64
tests/test_persist_terminal_fallback.py
Normal file
64
tests/test_persist_terminal_fallback.py
Normal file
|
|
@ -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",
|
||||
)
|
||||
46
tests/test_pg_compat.py
Normal file
46
tests/test_pg_compat.py
Normal file
|
|
@ -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"}
|
||||
78
tests/test_sa_admin_user_stats_tool_health.py
Normal file
78
tests/test_sa_admin_user_stats_tool_health.py
Normal file
|
|
@ -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
|
||||
41
tests/test_sa_app_settings.py
Normal file
41
tests/test_sa_app_settings.py
Normal file
|
|
@ -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
|
||||
132
tests/test_sa_app_users.py
Normal file
132
tests/test_sa_app_users.py
Normal file
|
|
@ -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
|
||||
75
tests/test_sa_auth_sessions.py
Normal file
75
tests/test_sa_auth_sessions.py
Normal file
|
|
@ -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
|
||||
91
tests/test_sa_chat_messages.py
Normal file
91
tests/test_sa_chat_messages.py
Normal file
|
|
@ -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"
|
||||
52
tests/test_sa_chat_sessions.py
Normal file
52
tests/test_sa_chat_sessions.py
Normal file
|
|
@ -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
|
||||
85
tests/test_sa_fork_trim_admin_sessions.py
Normal file
85
tests/test_sa_fork_trim_admin_sessions.py
Normal file
|
|
@ -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)
|
||||
54
tests/test_sa_prune_orphans.py
Normal file
54
tests/test_sa_prune_orphans.py
Normal file
|
|
@ -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)
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Add a link
Reference in a new issue