mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 03:30:48 +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
|
|
@ -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)
|
||||
48
tests/test_sa_tenant_bind_code.py
Normal file
48
tests/test_sa_tenant_bind_code.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
"""SA migration: tenant + bind_code."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from svc.persistence.db.engine import clear_assistant_engine_cache
|
||||
from svc.persistence.sqlite_store import SqliteStore
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path):
|
||||
monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False)
|
||||
monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path))
|
||||
dbfile = tmp_path / "sa_tenant_bc.sqlite"
|
||||
monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile))
|
||||
clear_assistant_engine_cache()
|
||||
s = SqliteStore(str(dbfile))
|
||||
try:
|
||||
yield s
|
||||
finally:
|
||||
clear_assistant_engine_cache()
|
||||
|
||||
|
||||
def test_sa_tenant_crud(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
t = s.create_tenant("Acme")
|
||||
listed = s.list_tenants(limit=50)
|
||||
assert any(x["id"] == t["id"] for x in listed)
|
||||
assert s.delete_tenant(tenant_id=t["id"]) == 1
|
||||
assert s.delete_tenant(tenant_id=t["id"]) == 0
|
||||
|
||||
|
||||
def test_sa_bind_code_list_and_consume(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
t = s.create_tenant("BindCo")
|
||||
s.create_bind_code(tenant_id=t["id"], role="member", code="CODE12345")
|
||||
rows = s.list_bind_codes(tenant_id=t["id"], limit=10)
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["code"] == "CODE12345"
|
||||
assert rows[0]["used_at"] is None
|
||||
out = s.consume_bind_code(code="CODE12345", channel="wecom", external_user_id="wx-1", display_name="Ext")
|
||||
assert out is not None
|
||||
assert out["tenant_id"] == t["id"]
|
||||
rows2 = s.list_bind_codes(tenant_id=t["id"], limit=10)
|
||||
assert rows2[0]["used_at"] is not None
|
||||
assert rows2[0]["used_by_external_user_id"] == "wx-1"
|
||||
assert s.consume_bind_code(code="CODE12345", channel="wecom", external_user_id="wx-2") is None
|
||||
121
tests/test_sa_tool_log_trace.py
Normal file
121
tests/test_sa_tool_log_trace.py
Normal file
|
|
@ -0,0 +1,121 @@
|
|||
"""SA migration: tool_log MCP queries + trace_event writes/lists."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from svc.persistence.db.engine import clear_assistant_engine_cache
|
||||
from svc.persistence.sqlite_store import SqliteStore
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path):
|
||||
monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False)
|
||||
monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path))
|
||||
dbfile = tmp_path / "sa_tl_tr.sqlite"
|
||||
monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile))
|
||||
clear_assistant_engine_cache()
|
||||
s = SqliteStore(str(dbfile))
|
||||
try:
|
||||
yield s
|
||||
finally:
|
||||
clear_assistant_engine_cache()
|
||||
|
||||
|
||||
def test_sa_mcp_tool_summaries_and_call_logs(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
sess = s.create_session("MCP2")
|
||||
s.add_tool_log(sess.id, "mcp__srv99__x", {}, {"ok": True}, specialist="sp1")
|
||||
s.add_tool_log(sess.id, "mcp__srv99__x", {}, {"ok": True}, specialist="sp1")
|
||||
agg = s.list_mcp_tool_aggregate_usage()
|
||||
assert "mcp__srv99__x" in agg
|
||||
assert agg["mcp__srv99__x"]["count"] == 2
|
||||
summ = s.list_mcp_tool_usage_summary(limit=50)
|
||||
hit = next(x for x in summ if x["tool_name"] == "mcp__srv99__x")
|
||||
assert hit["count"] == 2
|
||||
assert hit["server_id"] == "srv99"
|
||||
logs = s.list_mcp_tool_call_logs(server_id="srv99", limit=10)
|
||||
assert len(logs) == 2
|
||||
assert logs[0]["session_id"] == sess.id
|
||||
|
||||
|
||||
def test_sa_trace_roundtrip_batch_and_window(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
sess = s.create_session("TR")
|
||||
s.add_trace_event(
|
||||
session_id=sess.id,
|
||||
trace_id="t1",
|
||||
span_id="s0",
|
||||
parent_span_id=None,
|
||||
event_type="noise",
|
||||
payload={"a": 1},
|
||||
)
|
||||
s.add_trace_events_batch(
|
||||
[
|
||||
{
|
||||
"session_id": sess.id,
|
||||
"trace_id": "t1",
|
||||
"span_id": "s1",
|
||||
"parent_span_id": None,
|
||||
"event_type": "turn_started",
|
||||
"payload": {},
|
||||
},
|
||||
{
|
||||
"session_id": sess.id,
|
||||
"trace_id": "t1",
|
||||
"span_id": "s2",
|
||||
"parent_span_id": None,
|
||||
"event_type": "turn_finished",
|
||||
"payload": {},
|
||||
},
|
||||
]
|
||||
)
|
||||
desc = s.list_trace_events(session_id=sess.id, limit=10)
|
||||
assert len(desc) == 3
|
||||
asc_rows = s.list_trace_events_for_trace(session_id=sess.id, trace_id="t1", limit=50)
|
||||
assert len(asc_rows) == 3
|
||||
assert asc_rows[0]["event_type"] == "noise"
|
||||
st, en = s.get_turn_time_window(session_id=sess.id, trace_id="t1")
|
||||
assert st is not None and en is not None
|
||||
|
||||
|
||||
def test_sa_add_tool_log_get_tool_logs(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
sess = s.create_session("TL")
|
||||
s.add_tool_log(sess.id, "my_tool", {"a": 1}, {"ok": True}, specialist="spec", duration_ms=42)
|
||||
logs = s.get_tool_logs(sess.id, limit=10)
|
||||
assert len(logs) == 1
|
||||
assert logs[0]["tool_name"] == "my_tool"
|
||||
assert logs[0]["specialist"] == "spec"
|
||||
assert logs[0]["args"] == {"a": 1}
|
||||
assert logs[0]["result"] == {"ok": True}
|
||||
assert logs[0]["duration_ms"] == 42
|
||||
|
||||
|
||||
def test_sa_list_messages_in_time_window(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
sess = s.create_session("Win")
|
||||
t0 = "2024-01-01T10:00:00+00:00"
|
||||
t1 = "2024-01-01T10:05:00+00:00"
|
||||
t2 = "2024-01-01T10:10:00+00:00"
|
||||
s.add_message(sess.id, "user", "a", timestamp=t0)
|
||||
s.add_message(sess.id, "user", "b", timestamp=t1)
|
||||
s.add_message(sess.id, "user", "c", timestamp=t2)
|
||||
rows = s.list_messages_in_time_window(session_id=sess.id, start_ts=t0, end_ts=t1, limit=50)
|
||||
assert len(rows) == 2
|
||||
assert [x["content"] for x in rows] == ["a", "b"]
|
||||
|
||||
|
||||
def test_sa_move_tool_logs_to_session(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
src = s.create_session("SrcS")
|
||||
dst = s.create_session("DstS")
|
||||
s.add_tool_log(src.id, "t_move", {}, {"ok": 1})
|
||||
n = s.move_tool_logs_to_session(from_session_id=src.id, to_session_id=dst.id)
|
||||
assert n == 1
|
||||
assert s.get_tool_logs(src.id, limit=10) == []
|
||||
logs = s.get_tool_logs(dst.id, limit=10)
|
||||
assert len(logs) == 1
|
||||
assert logs[0]["tool_name"] == "t_move"
|
||||
assert s.move_tool_logs_to_session(from_session_id=src.id, to_session_id=dst.id) == 0
|
||||
assert s.move_tool_logs_to_session(from_session_id="", to_session_id=dst.id) == 0
|
||||
96
tests/test_sa_ui_session_owner.py
Normal file
96
tests/test_sa_ui_session_owner.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
"""SA migration: ui_session_owner upsert + backfill."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from svc.persistence.db.engine import clear_assistant_engine_cache
|
||||
from svc.persistence.sqlite_store import SqliteStore
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fresh_sqlite_store(monkeypatch: pytest.MonkeyPatch, tmp_path):
|
||||
monkeypatch.delenv("AIA_ASSISTANT_DB_BACKEND", raising=False)
|
||||
monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path))
|
||||
dbfile = tmp_path / "sa_ui_owner.sqlite"
|
||||
monkeypatch.setenv("AIA_ASSISTANT_DB_PATH", str(dbfile))
|
||||
clear_assistant_engine_cache()
|
||||
s = SqliteStore(str(dbfile))
|
||||
try:
|
||||
yield s
|
||||
finally:
|
||||
clear_assistant_engine_cache()
|
||||
|
||||
|
||||
def test_sa_create_session_for_user_and_get_owner(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
t = s.create_tenant("UIO")
|
||||
u = s.create_user(tenant_id=t["id"], display_name="U", role="member")
|
||||
sess = s.create_session_for_user(title="T", tenant_id=t["id"], user_id=u["id"])
|
||||
row = s.get_ui_session_owner(session_id=sess.id)
|
||||
assert row is not None
|
||||
assert row["tenant_id"] == t["id"]
|
||||
assert row["user_id"] == u["id"]
|
||||
s.ensure_ui_session_owner(session_id=sess.id, tenant_id=t["id"], user_id=u["id"])
|
||||
row2 = s.get_ui_session_owner(session_id=sess.id)
|
||||
assert row2 == row
|
||||
|
||||
|
||||
def test_sa_upsert_replace_changes_owner(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
t = s.create_tenant("UIO2")
|
||||
u1 = s.create_user(tenant_id=t["id"], display_name="A", role="member")
|
||||
u2 = s.create_user(tenant_id=t["id"], display_name="B", role="member")
|
||||
sess = s.create_session_for_user(title="X", tenant_id=t["id"], user_id=u1["id"])
|
||||
s._ui_session_owner_repo().upsert_replace(
|
||||
session_id=sess.id,
|
||||
tenant_id=t["id"],
|
||||
user_id=u2["id"],
|
||||
created_at="2099-01-01T00:00:00+00:00",
|
||||
)
|
||||
row = s.get_ui_session_owner(session_id=sess.id)
|
||||
assert row is not None
|
||||
assert row["user_id"] == u2["id"]
|
||||
|
||||
|
||||
def test_sa_backfill_orphan_sessions(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
t = s.create_tenant("BF")
|
||||
u = s.create_user(tenant_id=t["id"], display_name="BFU", role="member")
|
||||
orphan = s.create_session("orphan title")
|
||||
assert s.get_ui_session_owner(session_id=orphan.id) is None
|
||||
n = s.backfill_orphan_chat_sessions_for_user(tenant_id=t["id"], user_id=u["id"])
|
||||
assert n >= 1
|
||||
row = s.get_ui_session_owner(session_id=orphan.id)
|
||||
assert row is not None
|
||||
assert row["user_id"] == u["id"]
|
||||
|
||||
|
||||
def test_sa_backfill_ui_session_owner_from_channel_v2(fresh_sqlite_store: SqliteStore) -> None:
|
||||
s = fresh_sqlite_store
|
||||
t = s.create_tenant("CH")
|
||||
u = s.create_user(tenant_id=t["id"], display_name="CHU", role="member")
|
||||
sess = s.create_session("ch sess")
|
||||
ts = "2025-01-01T00:00:00+00:00"
|
||||
with s._connect() as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO channel_identity_v2
|
||||
(tenant_id, channel, account_id, external_user_id, user_id, created_at)
|
||||
VALUES (?, 'wecom', 'acc1', 'extu', ?, ?)
|
||||
""",
|
||||
(t["id"], u["id"], ts),
|
||||
)
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO channel_session_v2
|
||||
(tenant_id, channel, account_id, external_chat_id, external_user_id, session_id, created_at)
|
||||
VALUES (?, 'wecom', 'acc1', 'chat1', 'extu', ?, ?)
|
||||
""",
|
||||
(t["id"], sess.id, ts),
|
||||
)
|
||||
n = s.backfill_ui_session_owner_from_channel_v2()
|
||||
assert n >= 1
|
||||
row = s.get_ui_session_owner(session_id=sess.id)
|
||||
assert row is not None
|
||||
assert row["user_id"] == u["id"]
|
||||
Loading…
Add table
Add a link
Reference in a new issue