diff --git a/bench/ecommerce_demo.py b/bench/ecommerce_demo.py index 0fab70d..33cdb88 100644 --- a/bench/ecommerce_demo.py +++ b/bench/ecommerce_demo.py @@ -28,7 +28,12 @@ from lang2sql.harness.loop import agent_loop from lang2sql.safety.pipeline import SafetyPipeline from lang2sql.tenancy.concierge import ContextConcierge -from lang2sql.tools.semantic_federation import FedEntry, _kv_key, _render_effective, _load_all, _resolve_term +from lang2sql.tools.semantic_federation import ( + FedEntry, + _kv_key, + _load_all, + _render_effective, +) # Stable IDs for the demo guild and its two channels. GUILD = "acme-shop" @@ -50,7 +55,9 @@ def _finance_identity() -> Identity: return Identity(user_id="evan", guild_id=GUILD, channel_id=CH_FINANCE) -def _define_term(store: SqliteStore, scope: str, term: str, layer: str, entity: str, definition: str) -> None: +def _define_term( + store: SqliteStore, scope: str, term: str, layer: str, entity: str, definition: str +) -> None: entry = FedEntry(term=term, layer=layer, entity=entity, definition=definition) store.kv_set(scope, _kv_key(term, layer, entity), entry.to_json()) @@ -93,7 +100,9 @@ async def section_1_define_metrics(store: SqliteStore) -> None: rendered = _render_effective(store, scope, channel_id, ident.user_id) lines = [l for l in rendered.splitlines() if l.startswith("-")] - print(f"\nEffective layer for #{CH_MARKETING} now holds {len(lines)} definition(s):") + print( + f"\nEffective layer for #{CH_MARKETING} now holds {len(lines)} definition(s):" + ) print(rendered) @@ -104,10 +113,22 @@ async def section_2_federation(store: SqliteStore) -> None: mkt = _marketing_identity() fin = _finance_identity() - _define_term(store, GUILD, "active_user", "channel", CH_MARKETING, - "user with a login event in the last 30 days") - _define_term(store, GUILD, "active_user", "channel", CH_FINANCE, - "user with an active paid subscription") + _define_term( + store, + GUILD, + "active_user", + "channel", + CH_MARKETING, + "user with a login event in the last 30 days", + ) + _define_term( + store, + GUILD, + "active_user", + "channel", + CH_FINANCE, + "user with an active paid subscription", + ) print("Defined 'active_user' independently in two channels.\n") print("Now resolving the *effective* definition each channel sees") @@ -127,9 +148,9 @@ async def section_2_federation(store: SqliteStore) -> None: print(f" #{CH_MARKETING:<10} active_user → {mkt_def}") print(f" #{CH_FINANCE:<10} active_user → {fin_def}") - assert mkt_def and fin_def and mkt_def != fin_def, ( - f"Federation failed: mkt_def={mkt_def!r}, fin_def={fin_def!r}" - ) + assert ( + mkt_def and fin_def and mkt_def != fin_def + ), f"Federation failed: mkt_def={mkt_def!r}, fin_def={fin_def!r}" print("\n ✅ Same term, two live definitions, zero conflict.") print(" Each channel is its own branch in the federation tree;") print(" neither overwrote the other. (Wren's single MDL cannot do this.)") diff --git a/pyproject.toml b/pyproject.toml index fa6c6f2..572d328 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,6 +15,7 @@ dependencies = [ "discord.py>=2.3,<3.0", # Phase 1 frontend transport "cryptography>=42.0", # EncryptedSecrets at-rest encryption "sqlalchemy>=2.0", # generic DB explorer (one adapter, many engines) + "PyYAML>=6.0", # OKF bundle frontmatter serialization ] [project.optional-dependencies] diff --git a/src/lang2sql/adapters/db/d1_explorer.py b/src/lang2sql/adapters/db/d1_explorer.py index 1eda630..a719e9b 100644 --- a/src/lang2sql/adapters/db/d1_explorer.py +++ b/src/lang2sql/adapters/db/d1_explorer.py @@ -43,7 +43,9 @@ def __init__( ) -> None: self.account_id = account_id self.database_id = database_id - self._token = token if token is not None else os.environ.get("CLOUDFLARE_API_TOKEN") + self._token = ( + token if token is not None else os.environ.get("CLOUDFLARE_API_TOKEN") + ) self._timeout = timeout self._transport = transport or self._http_transport @@ -59,7 +61,9 @@ async def list_tables(self) -> list[Table]: async def describe_table(self, name: str) -> Table: rows = await self._query(f"PRAGMA table_info({_ident(name)})") cols = [ - Column(name=r["name"], type=r["type"] or "", nullable=not bool(r["notnull"])) + Column( + name=r["name"], type=r["type"] or "", nullable=not bool(r["notnull"]) + ) for r in rows ] return Table(name=name, schema="", columns=cols) @@ -86,7 +90,9 @@ async def _query(self, sql: str, params: list | None = None) -> list[dict]: def _http_transport(self, sql: str, params: list) -> dict: if not self._token: - raise RuntimeError("CLOUDFLARE_API_TOKEN not set (D1 requires an API token)") + raise RuntimeError( + "CLOUDFLARE_API_TOKEN not set (D1 requires an API token)" + ) url = ( f"{_API_ROOT}/accounts/{self.account_id}" f"/d1/database/{self.database_id}/query" diff --git a/src/lang2sql/adapters/db/dsn_builder.py b/src/lang2sql/adapters/db/dsn_builder.py index 188ea4f..b723041 100644 --- a/src/lang2sql/adapters/db/dsn_builder.py +++ b/src/lang2sql/adapters/db/dsn_builder.py @@ -36,7 +36,9 @@ def _quote(s: str) -> str: return quote_plus(s, safe="") -def build_postgresql(*, host: str, port: str, database: str, user: str, password: str) -> ConnectionSpec: +def build_postgresql( + *, host: str, port: str, database: str, user: str, password: str +) -> ConnectionSpec: # User may paste a full URL (e.g. "host/db?sslmode=require") into the host field. # Extract just the hostname to avoid corrupting the assembled DSN. parsed = urlsplit("//" + host) @@ -47,7 +49,9 @@ def build_postgresql(*, host: str, port: str, database: str, user: str, password return ConnectionSpec(dsn=dsn, extras={}) -def build_mysql(*, host: str, port: str, database: str, user: str, password: str) -> ConnectionSpec: +def build_mysql( + *, host: str, port: str, database: str, user: str, password: str +) -> ConnectionSpec: p = int(port) if port else 3306 dsn = f"mysql+pymysql://{_quote(user)}:{_quote(password)}@{host}:{p}/{database}" return ConnectionSpec(dsn=dsn, extras={}) @@ -143,7 +147,9 @@ def assemble(db_type: str, fields: dict[str, str]) -> ConnectionSpec: # Filter to the expected kwargs (modal can hand stray keys safely). expected = {name for name, *_ in FIELD_SCHEMA[db_type]} cleaned = {k: (v or "").strip() for k, v in fields.items() if k in expected} - missing = [n for n, _, req, _ in FIELD_SCHEMA[db_type] if req and not cleaned.get(n)] + missing = [ + n for n, _, req, _ in FIELD_SCHEMA[db_type] if req and not cleaned.get(n) + ] if missing: raise ValueError(f"missing required fields: {', '.join(missing)}") return builder(**cleaned) diff --git a/src/lang2sql/adapters/db/factory.py b/src/lang2sql/adapters/db/factory.py index b2180a0..fcde5a7 100644 --- a/src/lang2sql/adapters/db/factory.py +++ b/src/lang2sql/adapters/db/factory.py @@ -57,7 +57,7 @@ def build_explorer( # Normalize bare postgresql:// → postgresql+psycopg:// (psycopg3 is installed). if scheme == "postgresql": - connection = "postgresql+psycopg" + connection[len("postgresql"):] + connection = "postgresql+psycopg" + connection[len("postgresql") :] # Anything else is assumed to be a SQLAlchemy URL (driver loaded lazily). return SqlAlchemyExplorer(connection, schema=schema) diff --git a/src/lang2sql/adapters/db/postgres_explorer.py b/src/lang2sql/adapters/db/postgres_explorer.py index 1aaa33f..cfe6b50 100644 --- a/src/lang2sql/adapters/db/postgres_explorer.py +++ b/src/lang2sql/adapters/db/postgres_explorer.py @@ -18,7 +18,9 @@ columns=[ Column("id", "integer", nullable=False, description="Primary key."), Column("amount", "numeric", nullable=False, description="Order total."), - Column("status", "text", description="pending | paid | shipped | cancelled."), + Column( + "status", "text", description="pending | paid | shipped | cancelled." + ), Column("created_at", "timestamptz", nullable=False), ], ), @@ -36,8 +38,18 @@ _SAMPLES: dict[str, list[dict]] = { "public.orders": [ - {"id": 1, "amount": 49.90, "status": "paid", "created_at": "2026-05-01T10:00:00Z"}, - {"id": 2, "amount": 12.00, "status": "pending", "created_at": "2026-05-02T14:30:00Z"}, + { + "id": 1, + "amount": 49.90, + "status": "paid", + "created_at": "2026-05-01T10:00:00Z", + }, + { + "id": 2, + "amount": 12.00, + "status": "pending", + "created_at": "2026-05-02T14:30:00Z", + }, ], "public.users": [ {"id": 1, "email": "alice@example.com", "created_at": "2026-04-20T08:00:00Z"}, diff --git a/src/lang2sql/adapters/db/sqlalchemy_explorer.py b/src/lang2sql/adapters/db/sqlalchemy_explorer.py index c7129d8..11fd3cd 100644 --- a/src/lang2sql/adapters/db/sqlalchemy_explorer.py +++ b/src/lang2sql/adapters/db/sqlalchemy_explorer.py @@ -61,7 +61,9 @@ def _list_tables_sync(self) -> list[Table]: default = insp.default_schema_name effective = self._schema or default # Omit schema when it's the connection default so SQL stays unqualified. - display_schema = "" if (not self._schema or self._schema == default) else effective + display_schema = ( + "" if (not self._schema or self._schema == default) else effective + ) return [ Table(name=t, schema=display_schema) for t in insp.get_table_names(schema=self._schema) diff --git a/src/lang2sql/adapters/llm/fake.py b/src/lang2sql/adapters/llm/fake.py index 59dd7a0..29d45df 100644 --- a/src/lang2sql/adapters/llm/fake.py +++ b/src/lang2sql/adapters/llm/fake.py @@ -50,7 +50,9 @@ async def complete( ) # No tools at all → just answer. - return Completion(content="(no tools available) Hello from FakeLLM.", finish_reason="stop") + return Completion( + content="(no tools available) Hello from FakeLLM.", finish_reason="stop" + ) def _demo_args(spec: ToolSpec) -> str: diff --git a/src/lang2sql/adapters/llm/openai_.py b/src/lang2sql/adapters/llm/openai_.py index df3a112..e604efe 100644 --- a/src/lang2sql/adapters/llm/openai_.py +++ b/src/lang2sql/adapters/llm/openai_.py @@ -83,7 +83,9 @@ def _post(self, payload: dict[str, Any]) -> dict[str, Any]: try: return json.loads(text) except (ValueError, TypeError) as exc: - raise RuntimeError(f"OpenAI returned non-JSON response: {text[:200]!r}") from exc + raise RuntimeError( + f"OpenAI returned non-JSON response: {text[:200]!r}" + ) from exc def _strip_thinking(text: str) -> str: diff --git a/src/lang2sql/adapters/storage/okf_bundle.py b/src/lang2sql/adapters/storage/okf_bundle.py new file mode 100644 index 0000000..cf8f72d --- /dev/null +++ b/src/lang2sql/adapters/storage/okf_bundle.py @@ -0,0 +1,174 @@ +"""OkfBundle — OKF(Open Knowledge Format) 기반 지식 번들 어댑터. + +KV 캐시(SqliteStore)와 양방향 sync: +- export: KV → 스코프별 .md 파일 (Git 영속, 사람이 읽을 수 있는 형태) +- import_: .md 파일 → KV (번들에서 런타임 캐시 복원) + +디렉토리 구조 (OKF SPEC §3): + / + ├── guild/ + │ ├── index.md + │ ├── metrics/active_user.md + │ ├── tables/orders.md + │ ├── rules/exclude_cancelled.md + │ ├── dimensions/customer_tier.md + │ └── misc/.md # kind 미지정 + └── channel:/ + ├── index.md + └── metrics/active_user.md + +각 .md 파일 형식 (OKF SPEC §4): + --- + type: Metric + title: active_user + description: "30일 내 로그인한 users" + tags: [growth, retention] + applies_to: users + synonyms: [활성화고객] + layer: guild + entity: "" + inferred: false + timestamp: 2026-07-18T... + --- + + (markdown body — definition 반복 또는 추가 설명) +""" + +from __future__ import annotations + +from datetime import datetime, timezone +from pathlib import Path +from typing import TYPE_CHECKING + +import yaml + +from ...tools.semantic_federation import ( + FedEntry, + _KV_PREFIX, + _kv_key, +) + +if TYPE_CHECKING: + from .sqlite_store import SqliteStore + +_KIND_FOLDER: dict[str, str] = { + "metric": "metrics", + "table": "tables", + "rule": "rules", + "dimension": "dimensions", +} +_RESERVED = {"index.md", "log.md"} + + +class OkfBundle: + """KV ↔ OKF .md 파일 양방향 sync 어댑터.""" + + def __init__(self, base_dir: str) -> None: + self.base_dir = Path(base_dir) + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + def export(self, store: "SqliteStore", kv_scope: str) -> int: + """KV에서 모든 FedEntry를 읽어 .md 파일로 저장. 저장된 파일 수 반환.""" + raw = store.kv_list_prefix(kv_scope, _KV_PREFIX + ":") + count = 0 + for _key, val in raw: + try: + entry = FedEntry.from_json(val) + except (ValueError, KeyError): + continue + path = self._concept_path(entry) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(_entry_to_md(entry), encoding="utf-8") + count += 1 + return count + + def import_(self, store: "SqliteStore", kv_scope: str) -> int: + """번들의 .md 파일을 읽어 KV로 복원. 로드된 항목 수 반환.""" + count = 0 + for md_file in self.base_dir.rglob("*.md"): + if md_file.name in _RESERVED: + continue + entry = _md_to_entry(md_file) + if entry is None: + continue + key = _kv_key(entry.term, entry.layer, entry.entity) + store.kv_set(kv_scope, key, entry.to_json()) + count += 1 + return count + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + + def _scope_dir(self, entry: FedEntry) -> Path: + label = "guild" if entry.layer == "guild" else f"{entry.layer}:{entry.entity}" + return self.base_dir / label + + def _concept_path(self, entry: FedEntry) -> Path: + folder = _KIND_FOLDER.get(entry.kind, "misc") + slug = entry.term.strip().lower().replace(" ", "_").replace(":", "-") + return self._scope_dir(entry) / folder / f"{slug}.md" + + +# ------------------------------------------------------------------ +# Serialization helpers (module-level for testability) +# ------------------------------------------------------------------ + + +def _entry_to_md(entry: FedEntry) -> str: + """FedEntry → OKF .md 문자열 (SPEC §4.1).""" + fm: dict = { + "type": entry.kind.capitalize() if entry.kind else "Concept", + "title": entry.term, + "description": entry.definition, + } + if entry.tags: + fm["tags"] = entry.tags + if entry.applies_to: + fm["applies_to"] = entry.applies_to + if entry.synonyms: + fm["synonyms"] = entry.synonyms + fm["layer"] = entry.layer + fm["entity"] = entry.entity + fm["inferred"] = entry.inferred + fm["timestamp"] = datetime.now(timezone.utc).isoformat() + + yaml_block = yaml.dump( + fm, allow_unicode=True, default_flow_style=False, sort_keys=False + ) + return f"---\n{yaml_block}---\n\n{entry.definition}\n" + + +def _md_to_entry(path: Path) -> FedEntry | None: + """OKF .md 파일 → FedEntry. 파싱 실패 시 None 반환.""" + text = path.read_text(encoding="utf-8") + if not text.startswith("---"): + return None + try: + end = text.index("---", 3) + except ValueError: + return None + try: + fm = yaml.safe_load(text[3:end]) + except yaml.YAMLError: + return None + if not isinstance(fm, dict): + return None + + raw_kind = str(fm.get("type", "")).lower() + kind = raw_kind if raw_kind in _KIND_FOLDER else "" + + return FedEntry( + term=str(fm.get("title", path.stem)), + layer=str(fm.get("layer", "guild")), + entity=str(fm.get("entity", "")), + definition=str(fm.get("description", "")), + synonyms=fm.get("synonyms") or [], + inferred=bool(fm.get("inferred", False)), + kind=kind, + applies_to=str(fm.get("applies_to", "")), + tags=fm.get("tags") or [], + ) diff --git a/src/lang2sql/adapters/storage/sqlite_store.py b/src/lang2sql/adapters/storage/sqlite_store.py index 559c28b..2afea95 100644 --- a/src/lang2sql/adapters/storage/sqlite_store.py +++ b/src/lang2sql/adapters/storage/sqlite_store.py @@ -35,8 +35,7 @@ def __init__(self, path: str = ":memory:") -> None: self._create_tables() def _create_tables(self) -> None: - self._conn.executescript( - """ + self._conn.executescript(""" CREATE TABLE IF NOT EXISTS audit ( id INTEGER PRIMARY KEY AUTOINCREMENT, actor TEXT NOT NULL, @@ -55,8 +54,7 @@ def _create_tables(self) -> None: value TEXT NOT NULL, PRIMARY KEY (scope, key) ); - """ - ) + """) self._conn.commit() def close(self) -> None: @@ -125,9 +123,7 @@ def kv_set(self, scope: str, key: str, value: str) -> None: self._conn.commit() def kv_delete(self, scope: str, key: str) -> None: - self._conn.execute( - "DELETE FROM kv WHERE scope = ? AND key = ?", (scope, key) - ) + self._conn.execute("DELETE FROM kv WHERE scope = ? AND key = ?", (scope, key)) self._conn.commit() @staticmethod diff --git a/src/lang2sql/core/__init__.py b/src/lang2sql/core/__init__.py index 4aabff6..7357960 100644 --- a/src/lang2sql/core/__init__.py +++ b/src/lang2sql/core/__init__.py @@ -11,6 +11,13 @@ ) __all__ = [ - "Identity", "Scope", "ScopeLevel", - "Completion", "Message", "Role", "ToolCall", "ToolResult", "ToolSpec", + "Identity", + "Scope", + "ScopeLevel", + "Completion", + "Message", + "Role", + "ToolCall", + "ToolResult", + "ToolSpec", ] diff --git a/src/lang2sql/core/ports/__init__.py b/src/lang2sql/core/ports/__init__.py index 9928dea..69ae683 100644 --- a/src/lang2sql/core/ports/__init__.py +++ b/src/lang2sql/core/ports/__init__.py @@ -30,13 +30,29 @@ from .tool import ToolPort __all__ = [ - "AuditEvent", "AuditPort", - "Column", "ExplorerPort", "Table", - "FrontendPort", "InboundMessage", "OutboundMessage", - "CandidateKind", "DocExtractorPort", "Document", "SemanticCandidate", "SourcePort", + "AuditEvent", + "AuditPort", + "Column", + "ExplorerPort", + "Table", + "FrontendPort", + "InboundMessage", + "OutboundMessage", + "CandidateKind", + "DocExtractorPort", + "Document", + "SemanticCandidate", + "SourcePort", "LLMPort", - "ExtractorPort", "Fact", "RecallPort", "StorePort", - "SafetyContext", "SafetyDecision", "SafetyLayerPort", "SafetyPipelinePort", "Verdict", + "ExtractorPort", + "Fact", + "RecallPort", + "StorePort", + "SafetyContext", + "SafetyDecision", + "SafetyLayerPort", + "SafetyPipelinePort", + "Verdict", "SecretsPort", "ScopeResolverPort", "SessionStorePort", diff --git a/src/lang2sql/core/ports/audit.py b/src/lang2sql/core/ports/audit.py index 63940d8..e39d9bb 100644 --- a/src/lang2sql/core/ports/audit.py +++ b/src/lang2sql/core/ports/audit.py @@ -12,17 +12,16 @@ @dataclass class AuditEvent: - actor: str # user_id - action: str # "run_sql" | "define_metric" | "ingest" | ... - scope: str # session/scope key + actor: str # user_id + action: str # "run_sql" | "define_metric" | "ingest" | ... + scope: str # session/scope key detail: dict[str, Any] = field(default_factory=dict) - ts: float = 0.0 # epoch seconds; filled by the store if 0 + ts: float = 0.0 # epoch seconds; filled by the store if 0 @runtime_checkable class AuditPort(Protocol): - async def record(self, event: AuditEvent) -> None: - ... + async def record(self, event: AuditEvent) -> None: ... async def query(self, actor: str, limit: int = 20) -> list[AuditEvent]: """Recent events for one actor, newest first.""" diff --git a/src/lang2sql/core/ports/frontend.py b/src/lang2sql/core/ports/frontend.py index c56df5e..f9d6049 100644 --- a/src/lang2sql/core/ports/frontend.py +++ b/src/lang2sql/core/ports/frontend.py @@ -28,7 +28,7 @@ class OutboundMessage: """Normalised agent output a frontend renders natively.""" text: str - file_bytes: bytes | None = None # e.g. CSV when result > 50 rows + file_bytes: bytes | None = None # e.g. CSV when result > 50 rows file_name: str | None = None diff --git a/src/lang2sql/core/ports/ingestion.py b/src/lang2sql/core/ports/ingestion.py index 38f37fe..899a38d 100644 --- a/src/lang2sql/core/ports/ingestion.py +++ b/src/lang2sql/core/ports/ingestion.py @@ -19,7 +19,7 @@ class Document: name: str text: str - source_id: str = "" # preserved on resulting semantic entries + source_id: str = "" # preserved on resulting semantic entries class CandidateKind(str, Enum): @@ -43,13 +43,11 @@ class SemanticCandidate: class SourcePort(Protocol): """Where a document comes from (file/URL/Notion/…). Axis 1.""" - async def fetch(self, ref: str, blob: bytes | None = None) -> Document: - ... + async def fetch(self, ref: str, blob: bytes | None = None) -> Document: ... @runtime_checkable class DocExtractorPort(Protocol): """How definitions are pulled from a document (LLM/DDL/…). Axis 2.""" - async def extract(self, doc: Document) -> list[SemanticCandidate]: - ... + async def extract(self, doc: Document) -> list[SemanticCandidate]: ... diff --git a/src/lang2sql/core/ports/memory.py b/src/lang2sql/core/ports/memory.py index 14ffbc6..7d029ec 100644 --- a/src/lang2sql/core/ports/memory.py +++ b/src/lang2sql/core/ports/memory.py @@ -18,7 +18,7 @@ class Fact: """A remembered statement, scoped to a user/conversation.""" id: str - owner: str # user_id or scope key + owner: str # user_id or scope key text: str source: str = "manual" # "manual" (/remember) | "auto" (v1.5 extractor) ts: float = 0.0 @@ -40,8 +40,7 @@ class RecallPort(Protocol): V1 returns everything; v1.5 filters by keyword, v2 by vector similarity. """ - async def recall(self, owner: str, query: str, store: StorePort) -> list[Fact]: - ... + async def recall(self, owner: str, query: str, store: StorePort) -> list[Fact]: ... @runtime_checkable @@ -52,5 +51,6 @@ class ExtractorPort(Protocol): yields nothing); v1.5 mines the transcript with an LLM. """ - async def extract(self, owner: str, transcript: Sequence[Message]) -> list[Fact]: - ... + async def extract( + self, owner: str, transcript: Sequence[Message] + ) -> list[Fact]: ... diff --git a/src/lang2sql/core/ports/safety.py b/src/lang2sql/core/ports/safety.py index 8b69d11..52711aa 100644 --- a/src/lang2sql/core/ports/safety.py +++ b/src/lang2sql/core/ports/safety.py @@ -16,17 +16,17 @@ class Verdict(str, Enum): PASS = "pass" BLOCK = "block" - CONFIRM = "confirm" # ask the user before proceeding - REWRITE = "rewrite" # layer rewrote the SQL (e.g. attach LIMIT) + CONFIRM = "confirm" # ask the user before proceeding + REWRITE = "rewrite" # layer rewrote the SQL (e.g. attach LIMIT) @dataclass class SafetyDecision: verdict: Verdict - sql: str # possibly rewritten + sql: str # possibly rewritten reason: str = "" - layer: str = "" # which layer decided - confirm_prompt: str = "" # populated when verdict is CONFIRM + layer: str = "" # which layer decided + confirm_prompt: str = "" # populated when verdict is CONFIRM @dataclass @@ -45,16 +45,14 @@ class SafetyLayerPort(Protocol): @property def name(self) -> str: ... - def check(self, sql: str, ctx: SafetyContext) -> SafetyDecision: - ... + def check(self, sql: str, ctx: SafetyContext) -> SafetyDecision: ... @runtime_checkable class SafetyPipelinePort(Protocol): """Runs layers in order; first non-PASS short-circuits.""" - def evaluate(self, sql: str, ctx: SafetyContext) -> SafetyDecision: - ... + def evaluate(self, sql: str, ctx: SafetyContext) -> SafetyDecision: ... @property def layers(self) -> Sequence[SafetyLayerPort]: ... diff --git a/src/lang2sql/core/ports/secrets.py b/src/lang2sql/core/ports/secrets.py index 92c7fda..379b5aa 100644 --- a/src/lang2sql/core/ports/secrets.py +++ b/src/lang2sql/core/ports/secrets.py @@ -20,5 +20,4 @@ async def set(self, scope: str, key: str, value: str) -> None: """Encrypt and persist one secret under ``scope``.""" ... - async def delete(self, scope: str, key: str) -> None: - ... + async def delete(self, scope: str, key: str) -> None: ... diff --git a/src/lang2sql/core/ports/session_store.py b/src/lang2sql/core/ports/session_store.py index 5a3a763..437fdd8 100644 --- a/src/lang2sql/core/ports/session_store.py +++ b/src/lang2sql/core/ports/session_store.py @@ -18,5 +18,4 @@ async def load(self, key: str) -> "Session | None": """Restore a saved session, or ``None`` for a fresh conversation.""" ... - async def save(self, key: str, session: "Session") -> None: - ... + async def save(self, key: str, session: "Session") -> None: ... diff --git a/src/lang2sql/frontends/discord/bot.py b/src/lang2sql/frontends/discord/bot.py index 84cf982..29915f2 100644 --- a/src/lang2sql/frontends/discord/bot.py +++ b/src/lang2sql/frontends/discord/bot.py @@ -132,31 +132,60 @@ def _register_commands(self) -> None: tree = self.tree handlers = self._handlers - @tree.command(name="setup", description="Connect a database with a guided form (no DSN needed)") + @tree.command( + name="setup", + description="Connect a database with a guided form (no DSN needed)", + ) async def setup(interaction: discord.Interaction) -> None: - from .setup_wizard import start_setup_flow # local import — discord-only path + from .setup_wizard import ( + start_setup_flow, + ) # local import — discord-only path + await start_setup_flow(interaction, handlers, _interaction_context) @tree.command(name="connect", description="Store a database connection string") async def connect(interaction: discord.Interaction, dsn: str) -> None: - await self._run(interaction, handlers.connect(to_identity(_interaction_context(interaction)), dsn)) + await self._run( + interaction, + handlers.connect(to_identity(_interaction_context(interaction)), dsn), + ) @tree.command(name="ingest", description="Propose definitions from a document") async def ingest(interaction: discord.Interaction, ref: str) -> None: - await self._run(interaction, handlers.ingest(to_identity(_interaction_context(interaction)), ref=ref)) + await self._run( + interaction, + handlers.ingest( + to_identity(_interaction_context(interaction)), ref=ref + ), + ) @tree.command(name="remember", description="Remember a fact for future turns") async def remember(interaction: discord.Interaction, text: str) -> None: - await self._run(interaction, handlers.remember(to_identity(_interaction_context(interaction)), text)) + await self._run( + interaction, + handlers.remember(to_identity(_interaction_context(interaction)), text), + ) - @tree.command(name="enrich", description="LLM으로 DB 컬럼 메타데이터 자동 보강 (clear=True로 초기화)") - async def enrich(interaction: discord.Interaction, table: str = "", clear: bool = False) -> None: + @tree.command( + name="enrich", + description="LLM으로 DB 컬럼 메타데이터 자동 보강 (clear=True로 초기화)", + ) + async def enrich( + interaction: discord.Interaction, table: str = "", clear: bool = False + ) -> None: await self._run( interaction, - handlers.enrich(to_identity(_interaction_context(interaction)), table=table, clear=clear), + handlers.enrich( + to_identity(_interaction_context(interaction)), + table=table, + clear=clear, + ), ) - @tree.command(name="term_custom", description="비즈니스 용어 등록·조회·삭제 (action: show / remove, term: 용어명)") + @tree.command( + name="term_custom", + description="비즈니스 용어 등록·조회·삭제 (action: show / remove, term: 용어명)", + ) async def term_custom( interaction: discord.Interaction, action: str = "", @@ -167,12 +196,19 @@ async def term_custom( if action == "show": await self._run(interaction, handlers.term_custom(ident, list_all=True)) elif action == "remove": - await self._run(interaction, handlers.term_custom(ident, term=term, layer=layer, remove=True)) + await self._run( + interaction, + handlers.term_custom(ident, term=term, layer=layer, remove=True), + ) else: from .term_wizard import start_term_add_flow + await start_term_add_flow(interaction, handlers, _interaction_context) - @tree.command(name="org_setup", description="조직(전사) 또는 팀(채널) 등록 + DB 스캔으로 비즈니스 용어 자동 추출") + @tree.command( + name="org_setup", + description="조직(전사) 또는 팀(채널) 등록 + DB 스캔으로 비즈니스 용어 자동 추출", + ) async def org_setup( interaction: discord.Interaction, org: str = "", @@ -181,12 +217,20 @@ async def org_setup( ) -> None: await self._run( interaction, - handlers.org_setup(to_identity(_interaction_context(interaction)), org=org, team=team, clear=clear), + handlers.org_setup( + to_identity(_interaction_context(interaction)), + org=org, + team=team, + clear=clear, + ), ) @tree.command(name="audit_me", description="Show your recent activity") async def audit_me(interaction: discord.Interaction) -> None: - await self._run(interaction, handlers.audit_me(to_identity(_interaction_context(interaction)))) + await self._run( + interaction, + handlers.audit_me(to_identity(_interaction_context(interaction))), + ) async def _run(self, interaction: discord.Interaction, coro) -> None: """Await a handler coroutine and reply with its OutboundMessage.""" @@ -197,9 +241,12 @@ async def _run(self, interaction: discord.Interaction, coro) -> None: await interaction.followup.send(**kwargs) except Exception as exc: import traceback + traceback.print_exc() try: - await interaction.followup.send(content=f"❌ Error: {type(exc).__name__}: {exc}") + await interaction.followup.send( + content=f"❌ Error: {type(exc).__name__}: {exc}" + ) except Exception: pass @@ -225,6 +272,7 @@ async def on_message(self, message: discord.Message) -> None: await message.channel.send(**kwargs) except Exception as exc: import traceback + traceback.print_exc() await message.channel.send(content=f"❌ Error: {type(exc).__name__}: {exc}") diff --git a/src/lang2sql/frontends/discord/commands.py b/src/lang2sql/frontends/discord/commands.py index fbd7f5b..4132706 100644 --- a/src/lang2sql/frontends/discord/commands.py +++ b/src/lang2sql/frontends/discord/commands.py @@ -79,7 +79,9 @@ async def query(self, identity: Identity, text: str) -> OutboundMessage: async def remember(self, identity: Identity, text: str) -> OutboundMessage: """Persist a user fact via the memory service (manual ``/remember``).""" ctx = await self._concierge.build_context(identity) - result = await ctx.tools.dispatch("remember", {"text": text}, ctx, "cmd:remember") + result = await ctx.tools.dispatch( + "remember", {"text": text}, ctx, "cmd:remember" + ) return OutboundMessage(text=result.content) async def audit_me(self, identity: Identity) -> OutboundMessage: @@ -149,7 +151,9 @@ async def register_db_for_guild( ) ) - async def enrich(self, identity: Identity, table: str = "", clear: bool = False) -> OutboundMessage: + async def enrich( + self, identity: Identity, table: str = "", clear: bool = False + ) -> OutboundMessage: """Run EnrichSchema tool: sample DB columns and LLM-infer descriptions.""" ctx = await self._concierge.build_context(identity) result = await ctx.tools.dispatch( @@ -163,7 +167,10 @@ async def org_setup( """조직(전사) 또는 팀(채널) 등록 + DB 스캔으로 비즈니스 용어 자동 추출.""" ctx = await self._concierge.build_context(identity) result = await ctx.tools.dispatch( - "org_setup", {"org": org, "team": team, "clear": clear}, ctx, "cmd:org_setup" + "org_setup", + {"org": org, "team": team, "clear": clear}, + ctx, + "cmd:org_setup", ) return OutboundMessage(text=result.content) @@ -184,9 +191,14 @@ async def term_custom( result = await ctx.tools.dispatch( "term_custom", { - "term": term, "definition": definition, "layer": layer, - "synonyms": synonyms, "inferred": inferred, "scan": scan, - "remove": remove, "list": list_all, + "term": term, + "definition": definition, + "layer": layer, + "synonyms": synonyms, + "inferred": inferred, + "scan": scan, + "remove": remove, + "list": list_all, }, ctx, "cmd:term_custom", diff --git a/src/lang2sql/frontends/discord/render.py b/src/lang2sql/frontends/discord/render.py index f143d2a..5f90be2 100644 --- a/src/lang2sql/frontends/discord/render.py +++ b/src/lang2sql/frontends/discord/render.py @@ -67,9 +67,7 @@ def render_answer( return OutboundMessage(text=text) -def _rows_to_csv( - rows: Sequence[Sequence[Any]], header: Sequence[str] | None -) -> str: +def _rows_to_csv(rows: Sequence[Sequence[Any]], header: Sequence[str] | None) -> str: """Serialise ``rows`` (optionally with a ``header``) to a CSV string.""" buf = io.StringIO() writer = csv.writer(buf) diff --git a/src/lang2sql/harness/context.py b/src/lang2sql/harness/context.py index d9af9b1..a9fdb8c 100644 --- a/src/lang2sql/harness/context.py +++ b/src/lang2sql/harness/context.py @@ -33,4 +33,5 @@ class HarnessContext: safety: SafetyPipelinePort | None = None audit: AuditPort | None = None store: SqliteStore | None = None + okf_bundle_dir: str | None = None max_turns: int = 8 diff --git a/src/lang2sql/harness/session.py b/src/lang2sql/harness/session.py index b3f31d6..3cd610b 100644 --- a/src/lang2sql/harness/session.py +++ b/src/lang2sql/harness/session.py @@ -30,12 +30,15 @@ def reset(self) -> None: def compress(self) -> None: """Remove tool call/result messages to prevent context pollution across turns.""" from ..core.types import Role + cleaned: list[Message] = [] for msg in self.transcript: if msg.role == Role.TOOL: continue if msg.role == Role.ASSISTANT and msg.tool_calls: - if msg.content: # skip if no text content — empty assistant messages confuse OpenAI + if ( + msg.content + ): # skip if no text content — empty assistant messages confuse OpenAI cleaned.append(Message(role=Role.ASSISTANT, content=msg.content)) else: cleaned.append(msg) diff --git a/src/lang2sql/harness/system_prompt.py b/src/lang2sql/harness/system_prompt.py index 31e70c4..aec656a 100644 --- a/src/lang2sql/harness/system_prompt.py +++ b/src/lang2sql/harness/system_prompt.py @@ -32,8 +32,7 @@ async def build_system_prompt(ctx: HarnessContext) -> str: if tables: scope = ctx.identity.kv_scope if ctx.store else None has_enrichment = bool( - scope and ctx.store and - ctx.store.kv_get(scope, "schema_relationships") + scope and ctx.store and ctx.store.kv_get(scope, "schema_relationships") ) if has_enrichment and scope and ctx.store: @@ -46,10 +45,19 @@ async def build_system_prompt(ctx: HarnessContext) -> str: continue col_lines = [] for col in described.columns: - desc = col.description or ctx.store.kv_get(scope, f"enriched_desc:{tbl.name}:{col.name}") or "" + desc = ( + col.description + or ctx.store.kv_get( + scope, f"enriched_desc:{tbl.name}:{col.name}" + ) + or "" + ) col_lines.append(f" - {col.name}{': ' + desc if desc else ''}") schema_lines.append(f"- {tbl.qualified}\n" + "\n".join(col_lines)) - parts.append("## Known tables (with column descriptions)\n" + "\n".join(schema_lines)) + parts.append( + "## Known tables (with column descriptions)\n" + + "\n".join(schema_lines) + ) else: names = ", ".join(t.qualified for t in tables) parts.append("## Known tables\n" + names) @@ -62,11 +70,14 @@ async def build_system_prompt(ctx: HarnessContext) -> str: rels = json.loads(raw) if rels: rel_text = "\n".join(f"- {r}" for r in rels) - parts.append("## Table relationships (use these for JOINs)\n" + rel_text) + parts.append( + "## Table relationships (use these for JOINs)\n" + rel_text + ) except (ValueError, TypeError): pass from ..tools.semantic_federation import build_prompt_section + user_id = ctx.identity.user_id or "unknown" channel_id = ctx.identity.effective_channel_id semfed_section = build_prompt_section(ctx.store, scope, channel_id, user_id) diff --git a/src/lang2sql/harness/tool_registry.py b/src/lang2sql/harness/tool_registry.py index c6add7e..b0122d7 100644 --- a/src/lang2sql/harness/tool_registry.py +++ b/src/lang2sql/harness/tool_registry.py @@ -28,10 +28,14 @@ async def dispatch( ) -> ToolResult: tool = self._tools.get(name) if tool is None: - return ToolResult(call_id=call_id, content=f"unknown tool: {name}", is_error=True) + return ToolResult( + call_id=call_id, content=f"unknown tool: {name}", is_error=True + ) try: result = await tool.run(args, ctx) result.call_id = call_id # tools don't know their call id; stamp it here return result except Exception as exc: # tools must never crash the loop - return ToolResult(call_id=call_id, content=f"{type(exc).__name__}: {exc}", is_error=True) + return ToolResult( + call_id=call_id, content=f"{type(exc).__name__}: {exc}", is_error=True + ) diff --git a/src/lang2sql/safety/layers/whitelist.py b/src/lang2sql/safety/layers/whitelist.py index 9a6a284..11ab87e 100644 --- a/src/lang2sql/safety/layers/whitelist.py +++ b/src/lang2sql/safety/layers/whitelist.py @@ -110,7 +110,9 @@ def check(self, sql: str, ctx: SafetyContext) -> SafetyDecision: else: # Consume leading bare option words (ANALYZE, VERBOSE, ...). while True: - m = re.match(r"(?i)^(ANALYZE|VERBOSE|COSTS|BUFFERS)\b\s*(.*)$", body) + m = re.match( + r"(?i)^(ANALYZE|VERBOSE|COSTS|BUFFERS)\b\s*(.*)$", body + ) if m is None: break body = m.group(2).strip() diff --git a/src/lang2sql/tenancy/concierge.py b/src/lang2sql/tenancy/concierge.py index 5812865..700116f 100644 --- a/src/lang2sql/tenancy/concierge.py +++ b/src/lang2sql/tenancy/concierge.py @@ -55,7 +55,9 @@ def __init__( ) -> None: self._store = store if store is not None else SqliteStore(path) self._llm = llm if llm is not None else _default_llm() - self._explorer = explorer or explorer_from_env() or PostgresExplorer(_DEFAULT_DSN) + self._explorer = ( + explorer or explorer_from_env() or PostgresExplorer(_DEFAULT_DSN) + ) self._safety = safety if safety is not None else SafetyPipeline() self._secrets = ( secrets if secrets is not None else EncryptedSecrets(self._store) @@ -64,7 +66,9 @@ def __init__( self._max_turns = max_turns # V1 memory (in-memory + inject-all + manual) and ingestion (file × LLM). - self._memory = MemoryService(InMemoryStore(), InjectAllRecall(), ManualExtractor()) + self._memory = MemoryService( + InMemoryStore(), InjectAllRecall(), ManualExtractor() + ) self._ingestion = IngestionPipeline() self._source = FileSource() self._extractor = LLMExtractor(self._llm) @@ -134,6 +138,7 @@ async def build_context( safety=self._safety, audit=self._audit, store=self._store, + okf_bundle_dir=os.getenv("OKF_BUNDLE_DIR"), max_turns=self._max_turns, ) diff --git a/src/lang2sql/tools/__init__.py b/src/lang2sql/tools/__init__.py index b8d0813..afc6dd6 100644 --- a/src/lang2sql/tools/__init__.py +++ b/src/lang2sql/tools/__init__.py @@ -13,6 +13,7 @@ from ..ingestion.pipeline import IngestionPipeline from ..memory.service import MemoryService from .ask_user import AskUser +from .confirm_ingest import ConfirmIngest from .enrich_schema import EnrichSchema from .explore_schema import ExploreSchema from .ingest_doc import IngestDoc @@ -23,8 +24,15 @@ __all__ = [ "build_default_tools", - "RunSQL", "ExploreSchema", "EnrichSchema", "SemanticFederationTool", - "OrgSetupTool", "Remember", "AskUser", "IngestDoc", + "RunSQL", + "ExploreSchema", + "EnrichSchema", + "SemanticFederationTool", + "OrgSetupTool", + "Remember", + "AskUser", + "IngestDoc", + "ConfirmIngest", ] @@ -45,4 +53,5 @@ def build_default_tools( AskUser(), Remember(memory), IngestDoc(ingestion, source, extractor), + ConfirmIngest(), ] diff --git a/src/lang2sql/tools/confirm_ingest.py b/src/lang2sql/tools/confirm_ingest.py new file mode 100644 index 0000000..339aa85 --- /dev/null +++ b/src/lang2sql/tools/confirm_ingest.py @@ -0,0 +1,164 @@ +"""confirm_ingest — 사용자가 승인한 ingest 후보를 시멘틱 레이어에 등록. + +ingest_doc이 KV에 저장한 pending_ingest:{ref} 후보 목록을 읽어 +FedEntry로 변환하고 KV에 저장한다. OKF_BUNDLE_DIR 환경변수가 설정된 경우 +OkfBundle로도 내보낸다. +""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING, Any + +from ..core.ports.ingestion import CandidateKind, SemanticCandidate +from ..core.types import ToolResult, ToolSpec +from ..tools.ingest_doc import PENDING_PREFIX +from ..tools.semantic_federation import FedEntry, _kv_key + +if TYPE_CHECKING: + from ..harness.context import HarnessContext + + +def _dict_to_candidate(d: dict) -> SemanticCandidate: + return SemanticCandidate( + kind=CandidateKind(d["kind"]), + name=d["name"], + definition=d["definition"], + applies_to=d.get("applies_to", ""), + source_id=d.get("source_id", ""), + ) + + +class ConfirmIngest: + """pending_ingest 후보를 KV(+OkfBundle)에 등록하는 툴.""" + + @property + def spec(self) -> ToolSpec: + return ToolSpec( + name="confirm_ingest", + description=( + "Register approved semantic candidates from a previously ingested document " + "into the semantic layer (KV store). " + "Call ingest_doc first to extract candidates." + ), + parameters={ + "type": "object", + "properties": { + "ref": { + "type": "string", + "description": "document ref used with ingest_doc", + }, + "accept": { + "type": "string", + "description": ( + "'all' to register every candidate, " + "or comma-separated 1-based indices like '1,3'" + ), + "default": "all", + }, + "layer": { + "type": "string", + "enum": ["guild", "channel", "member"], + "description": "scope to register under (default: channel)", + "default": "channel", + }, + }, + "required": ["ref"], + }, + ) + + async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: + ref = (args.get("ref") or "").strip() + accept = (args.get("accept") or "all").strip() + layer = (args.get("layer") or "channel").strip() + + if not ref: + return ToolResult(call_id="", content="'ref' is required.", is_error=True) + if ctx.store is None: + return ToolResult(call_id="", content="No store available.", is_error=True) + + kv_scope = ctx.identity.kv_scope + pending_key = f"{PENDING_PREFIX}:{ref}" + raw = ctx.store.kv_get(kv_scope, pending_key) + if not raw: + return ToolResult( + call_id="", + content=f"No pending candidates for '{ref}'. Run ingest_doc first.", + is_error=True, + ) + + try: + all_candidates = [_dict_to_candidate(d) for d in json.loads(raw)] + except (json.JSONDecodeError, KeyError, ValueError) as exc: + return ToolResult( + call_id="", + content=f"Pending data corrupted: {exc}", + is_error=True, + ) + + selected = _select(all_candidates, accept) + if selected is None: + return ToolResult( + call_id="", + content="Invalid accept value. Use 'all' or comma-separated indices like '1,3'.", + is_error=True, + ) + if not selected: + return ToolResult(call_id="", content="No candidates selected.") + + entity = _entity_for(ctx, layer) + registered: list[str] = [] + for cand in selected: + entry = FedEntry( + term=cand.name, + layer=layer, + entity=entity, + definition=cand.definition, + inferred=False, + kind=cand.kind.value, + applies_to=cand.applies_to, + ) + ctx.store.kv_set( + kv_scope, + _kv_key(entry.term, entry.layer, entry.entity), + entry.to_json(), + ) + registered.append(entry.term) + + ctx.store.kv_delete(kv_scope, pending_key) + + if ctx.okf_bundle_dir and registered: + from ..adapters.storage.okf_bundle import OkfBundle + + OkfBundle(ctx.okf_bundle_dir).export(ctx.store, kv_scope) + + kind_labels = {c.name: c.kind.value.upper() for c in selected} + lines = [f"✅ {len(registered)} term(s) registered to [{layer}]:"] + for t in registered: + lines.append(f" - [{kind_labels[t]}] {t}") + return ToolResult(call_id="", content="\n".join(lines)) + + +def _select( + candidates: list[SemanticCandidate], accept: str +) -> list[SemanticCandidate] | None: + if accept == "all": + return list(candidates) + try: + indices = [int(i.strip()) - 1 for i in accept.split(",") if i.strip()] + except ValueError: + return None + result = [] + for idx in indices: + if idx < 0 or idx >= len(candidates): + return None + result.append(candidates[idx]) + return result + + +def _entity_for(ctx: "HarnessContext", layer: str) -> str: + if layer == "channel": + return ctx.identity.effective_channel_id + if layer == "member": + return ctx.identity.user_id + return "" diff --git a/src/lang2sql/tools/enrich_schema.py b/src/lang2sql/tools/enrich_schema.py index f137398..c4ed4a7 100644 --- a/src/lang2sql/tools/enrich_schema.py +++ b/src/lang2sql/tools/enrich_schema.py @@ -85,23 +85,37 @@ def spec(self) -> ToolSpec: async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: if ctx.explorer is None: - return ToolResult(call_id="", content="DB가 연결되지 않았습니다 (/connect 먼저).", is_error=True) + return ToolResult( + call_id="", + content="DB가 연결되지 않았습니다 (/connect 먼저).", + is_error=True, + ) if ctx.store is None: - return ToolResult(call_id="", content="KV store를 사용할 수 없습니다.", is_error=True) + return ToolResult( + call_id="", content="KV store를 사용할 수 없습니다.", is_error=True + ) scope = ctx.identity.kv_scope if args.get("clear"): count = ctx.store.kv_delete_prefix(scope, _KV_PREFIX + ":") ctx.store.kv_delete(scope, _KV_RELATIONSHIPS) - return ToolResult(call_id="", content=f"🗑️ 보강 캐시 초기화 완료 ({count}개 삭제)") + return ToolResult( + call_id="", content=f"🗑️ 보강 캐시 초기화 완료 ({count}개 삭제)" + ) target = (args.get("table") or "").strip() all_tables = await ctx.explorer.list_tables() if target: - tables = [t for t in all_tables if t.name == target or t.qualified == target] + tables = [ + t for t in all_tables if t.name == target or t.qualified == target + ] if not tables: - return ToolResult(call_id="", content=f"테이블 '{target}'을 찾을 수 없습니다.", is_error=True) + return ToolResult( + call_id="", + content=f"테이블 '{target}'을 찾을 수 없습니다.", + is_error=True, + ) else: tables = all_tables @@ -117,7 +131,9 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: f"WHERE {col.name} IS NOT NULL LIMIT {_SAMPLE_LIMIT}" ) rows = await ctx.explorer.execute(sample_sql, _SAMPLE_LIMIT) - samples = [str(r.get(col.name, r.get(list(r.keys())[0], ""))) for r in rows] + samples = [ + str(r.get(col.name, r.get(list(r.keys())[0], ""))) for r in rows + ] except Exception: samples = [] sample_str = f" 샘플: {samples}" if samples else "" @@ -128,9 +144,7 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: prompt = _build_prompt(schema_block) # Single LLM call for all tables at once. - completion = await ctx.llm.complete( - [Message(role=Role.USER, content=prompt)] - ) + completion = await ctx.llm.complete([Message(role=Role.USER, content=prompt)]) columns, relationships = _extract_result(completion.content) if not columns and not relationships: @@ -153,7 +167,9 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: rel_lines: list[str] = [] if relationships: - ctx.store.kv_set(scope, _KV_RELATIONSHIPS, json.dumps(relationships, ensure_ascii=False)) + ctx.store.kv_set( + scope, _KV_RELATIONSHIPS, json.dumps(relationships, ensure_ascii=False) + ) rel_lines = [f"- {r}" for r in relationships] result_parts = [] @@ -162,4 +178,6 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: if rel_lines: result_parts.append("🔗 테이블 관계 추론:\n" + "\n".join(rel_lines)) - return ToolResult(call_id="", content="\n\n".join(result_parts) or "보강된 내용이 없습니다.") + return ToolResult( + call_id="", content="\n\n".join(result_parts) or "보강된 내용이 없습니다." + ) diff --git a/src/lang2sql/tools/explore_schema.py b/src/lang2sql/tools/explore_schema.py index 535266b..01e46e5 100644 --- a/src/lang2sql/tools/explore_schema.py +++ b/src/lang2sql/tools/explore_schema.py @@ -42,14 +42,19 @@ def spec(self) -> ToolSpec: parameters={ "type": "object", "properties": { - "table": {"type": "string", "description": "table name to describe; omit to list all tables"}, + "table": { + "type": "string", + "description": "table name to describe; omit to list all tables", + }, }, }, ) async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: if ctx.explorer is None: - return ToolResult(call_id="", content="no DB connected (use /connect)", is_error=True) + return ToolResult( + call_id="", content="no DB connected (use /connect)", is_error=True + ) table = (args.get("table") or "").strip() if not table: @@ -59,9 +64,12 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: t = await ctx.explorer.describe_table(table) t = _apply_enrichment_cache(t, ctx) - cols = "\n".join( - f"- {c.name}: {c.type}{'' if c.nullable else ' NOT NULL'}" - f"{(' — ' + c.description) if c.description else ''}" - for c in t.columns - ) or "(no columns)" + cols = ( + "\n".join( + f"- {c.name}: {c.type}{'' if c.nullable else ' NOT NULL'}" + f"{(' — ' + c.description) if c.description else ''}" + for c in t.columns + ) + or "(no columns)" + ) return ToolResult(call_id="", content=f"{t.qualified}\n{cols}") diff --git a/src/lang2sql/tools/ingest_doc.py b/src/lang2sql/tools/ingest_doc.py index e5118a6..85c573c 100644 --- a/src/lang2sql/tools/ingest_doc.py +++ b/src/lang2sql/tools/ingest_doc.py @@ -1,22 +1,35 @@ """ingest_doc — turn an uploaded document into semantic candidates (★③). Runs the Source × Extractor pipeline and returns the proposed metric/rule -definitions for the user to confirm. V1 does NOT auto-register — confirmation -is a frontend step (Discord buttons in Week 4); this tool surfaces the -candidates so the human stays in the loop (documents are the source of truth). +definitions for the user to confirm. Candidates are stored in KV under a +``pending_ingest:{ref}`` key so ``confirm_ingest`` can retrieve and register +them once the user approves. """ from __future__ import annotations +import json from typing import TYPE_CHECKING, Any -from ..core.ports.ingestion import DocExtractorPort, SourcePort +from ..core.ports.ingestion import DocExtractorPort, SemanticCandidate, SourcePort from ..core.types import ToolResult, ToolSpec from ..ingestion.pipeline import IngestionPipeline if TYPE_CHECKING: from ..harness.context import HarnessContext +PENDING_PREFIX = "pending_ingest" + + +def _candidate_to_dict(c: SemanticCandidate) -> dict: + return { + "kind": c.kind.value, + "name": c.name, + "definition": c.definition, + "applies_to": c.applies_to, + "source_id": c.source_id, + } + class IngestDoc: def __init__( @@ -35,30 +48,63 @@ def spec(self) -> ToolSpec: name="ingest_doc", description=( "Read a document and propose metric/dimension/rule definitions " - "for the user to confirm before they enter the semantic layer." + "for the user to confirm before they enter the semantic layer. " + "Use confirm_ingest to register the approved candidates." ), parameters={ "type": "object", "properties": { - "ref": {"type": "string", "description": "document path or identifier"}, - "content": {"type": "string", "description": "inline document text (alternative to ref)"}, + "ref": { + "type": "string", + "description": "document path or identifier", + }, + "content": { + "type": "string", + "description": "inline document text (alternative to ref)", + }, }, }, ) async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: - ref = (args.get("ref") or "inline").strip() + ref = (args.get("ref") or "").strip() content = args.get("content") blob = content.encode("utf-8") if isinstance(content, str) else None - if not content and ref == "inline": - return ToolResult(call_id="", content="provide a document 'ref' or inline 'content'", is_error=True) + if not content and not ref: + return ToolResult( + call_id="", + content="provide a document 'ref' or inline 'content'", + is_error=True, + ) + if not ref: + import hashlib + + ref = "inline:" + hashlib.md5(blob or b"").hexdigest()[:8] - candidates = await self._pipeline.ingest(self._source, self._extractor, ref, blob) + candidates = await self._pipeline.ingest( + self._source, self._extractor, ref, blob + ) if not candidates: - return ToolResult(call_id="", content="No definitions found in the document.") + return ToolResult( + call_id="", content="No definitions found in the document." + ) + + if ctx.store is not None: + pending_key = f"{PENDING_PREFIX}:{ref}" + ctx.store.kv_set( + ctx.identity.kv_scope, + pending_key, + json.dumps([_candidate_to_dict(c) for c in candidates]), + ) - lines = ["Proposed definitions (confirm to register):"] - for c in candidates: + lines = [f"Proposed definitions from '{ref}' (use confirm_ingest to register):"] + for i, c in enumerate(candidates, 1): applies = f" [{c.applies_to}]" if c.applies_to else "" - lines.append(f"- {c.kind.value.upper()} {c.name}{applies} → {c.definition}") + lines.append( + f" {i}. [{c.kind.value.upper()}] {c.name}{applies} — {c.definition}" + ) + lines.append( + f"\nRun confirm_ingest(ref='{ref}', accept='all') to register all, " + "or specify indices like accept='1,3'." + ) return ToolResult(call_id="", content="\n".join(lines)) diff --git a/src/lang2sql/tools/org_setup.py b/src/lang2sql/tools/org_setup.py index 4a4baf8..d981229 100644 --- a/src/lang2sql/tools/org_setup.py +++ b/src/lang2sql/tools/org_setup.py @@ -22,7 +22,12 @@ from ..core.ports.tool import ToolPort from ..core.types import Message, Role, ToolResult, ToolSpec -from .semantic_federation import FedEntry, _KV_PREFIX as _SEMFED_PREFIX, _kv_key as _semfed_kv_key, _parse_synonyms +from .semantic_federation import ( + FedEntry, + _KV_PREFIX as _SEMFED_PREFIX, + _kv_key as _semfed_kv_key, + _parse_synonyms, +) if TYPE_CHECKING: from ..harness.context import HarnessContext @@ -104,7 +109,11 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: team_name = str(args.get("team", "")).strip() if not org_name and not team_name: - return ToolResult(call_id="", content="❌ org 또는 team 파라미터가 필요합니다.", is_error=True) + return ToolResult( + call_id="", + content="❌ org 또는 team 파라미터가 필요합니다.", + is_error=True, + ) scope = ctx.identity.kv_scope channel_id = ctx.identity.effective_channel_id @@ -159,11 +168,17 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: ) if ctx.explorer is None: - return ToolResult(call_id="", content="❌ DB가 연결되지 않았습니다 (/setup 먼저).", is_error=True) + return ToolResult( + call_id="", + content="❌ DB가 연결되지 않았습니다 (/setup 먼저).", + is_error=True, + ) all_tables = await ctx.explorer.list_tables() if not all_tables: - return ToolResult(call_id="", content="❌ 접근 가능한 테이블이 없습니다.", is_error=True) + return ToolResult( + call_id="", content="❌ 접근 가능한 테이블이 없습니다.", is_error=True + ) schema_lines: list[str] = [] for tbl in all_tables: @@ -179,7 +194,9 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: f"WHERE {col.name} IS NOT NULL LIMIT {_SAMPLE_LIMIT}" ) rows = await ctx.explorer.execute(sample_sql, _SAMPLE_LIMIT) - samples = [str(r.get(col.name, r.get(list(r.keys())[0], ""))) for r in rows] + samples = [ + str(r.get(col.name, r.get(list(r.keys())[0], ""))) for r in rows + ] except Exception: samples = [] sample_str = f" 샘플: {samples}" if samples else "" @@ -202,7 +219,10 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: ctx.store.kv_set( scope, meta_key, - json.dumps({"name": display_name, "domain": domain, "registered_at": time.time()}, ensure_ascii=False), + json.dumps( + {"name": display_name, "domain": domain, "registered_at": time.time()}, + ensure_ascii=False, + ), ) saved_terms: list[str] = [] @@ -215,8 +235,12 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: if ":" in term: continue # `:` 포함 term은 KV 키 파싱을 깨트리므로 건너뜀 entry = FedEntry( - term=term, layer=layer, entity=entity, - definition=definition, synonyms=synonyms, inferred=True, + term=term, + layer=layer, + entity=entity, + definition=definition, + synonyms=synonyms, + inferred=True, ) kv_key = _semfed_kv_key(term, layer, entity) ctx.store.kv_set(scope, kv_key, entry.to_json()) diff --git a/src/lang2sql/tools/ping.py b/src/lang2sql/tools/ping.py index 3ded19c..0e65d0a 100644 --- a/src/lang2sql/tools/ping.py +++ b/src/lang2sql/tools/ping.py @@ -31,4 +31,6 @@ def spec(self) -> ToolSpec: async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: msg = args.get("message", "") - return ToolResult(call_id="", content=f"pong: {msg!r} (user={ctx.identity.user_id})") + return ToolResult( + call_id="", content=f"pong: {msg!r} (user={ctx.identity.user_id})" + ) diff --git a/src/lang2sql/tools/remember.py b/src/lang2sql/tools/remember.py index c48622f..09db60e 100644 --- a/src/lang2sql/tools/remember.py +++ b/src/lang2sql/tools/remember.py @@ -41,7 +41,11 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: fact = await self._memory.remember(ctx.identity.user_id, text) if ctx.audit is not None: await ctx.audit.record( - AuditEvent(actor=ctx.identity.user_id, action="remember", - scope=ctx.identity.session_key(), detail={"fact_id": fact.id}) + AuditEvent( + actor=ctx.identity.user_id, + action="remember", + scope=ctx.identity.session_key(), + detail={"fact_id": fact.id}, + ) ) return ToolResult(call_id="", content=f"🧠 Remembered: {text}") diff --git a/src/lang2sql/tools/run_sql.py b/src/lang2sql/tools/run_sql.py index 567f265..b348915 100644 --- a/src/lang2sql/tools/run_sql.py +++ b/src/lang2sql/tools/run_sql.py @@ -30,8 +30,14 @@ def spec(self) -> ToolSpec: parameters={ "type": "object", "properties": { - "sql": {"type": "string", "description": "a single SELECT or WITH query"}, - "limit": {"type": "integer", "description": "max rows (default 1000)"}, + "sql": { + "type": "string", + "description": "a single SELECT or WITH query", + }, + "limit": { + "type": "integer", + "description": "max rows (default 1000)", + }, }, "required": ["sql"], }, @@ -45,22 +51,40 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: limit = 1000 # tolerate a malformed limit from the model if ctx.safety is None: - return ToolResult(call_id="", content="run_sql unavailable: no safety pipeline wired", is_error=True) + return ToolResult( + call_id="", + content="run_sql unavailable: no safety pipeline wired", + is_error=True, + ) if ctx.explorer is None: - return ToolResult(call_id="", content="run_sql unavailable: no DB connected (use /connect)", is_error=True) + return ToolResult( + call_id="", + content="run_sql unavailable: no DB connected (use /connect)", + is_error=True, + ) decision = ctx.safety.evaluate(sql, SafetyContext(row_limit=limit)) if decision.verdict == Verdict.BLOCK: - return ToolResult(call_id="", content=f"BLOCKED by {decision.layer}: {decision.reason}", is_error=True) + return ToolResult( + call_id="", + content=f"BLOCKED by {decision.layer}: {decision.reason}", + is_error=True, + ) if decision.verdict == Verdict.CONFIRM: - return ToolResult(call_id="", content=f"NEEDS CONFIRMATION: {decision.confirm_prompt}") + return ToolResult( + call_id="", content=f"NEEDS CONFIRMATION: {decision.confirm_prompt}" + ) rows = await ctx.explorer.execute(decision.sql, limit) if ctx.audit is not None: await ctx.audit.record( - AuditEvent(actor=ctx.identity.user_id, action="run_sql", - scope=ctx.identity.session_key(), detail={"sql": decision.sql}) + AuditEvent( + actor=ctx.identity.user_id, + action="run_sql", + scope=ctx.identity.session_key(), + detail={"sql": decision.sql}, + ) ) return ToolResult(call_id="", content=_render_rows(decision.sql, rows)) diff --git a/src/lang2sql/tools/semantic_federation.py b/src/lang2sql/tools/semantic_federation.py index 6db15f1..169db4c 100644 --- a/src/lang2sql/tools/semantic_federation.py +++ b/src/lang2sql/tools/semantic_federation.py @@ -28,7 +28,10 @@ _KV_PREFIX = "cterm" _LAYERS = ("guild", "channel", "member") -from ..tools.enrich_schema import _KV_PREFIX as _ENRICH_PREFIX, _KV_RELATIONSHIPS as _ENRICH_RELATIONSHIPS +from ..tools.enrich_schema import ( + _KV_PREFIX as _ENRICH_PREFIX, + _KV_RELATIONSHIPS as _ENRICH_RELATIONSHIPS, +) _AMBIGUITY_SIGNALS: dict[str, str] = { r"(^|_)(created|registered|joined|signup)(_at|_date)?$": "신규/최초 가입 기준 용어", @@ -56,22 +59,33 @@ def _parse_synonyms(raw: Any) -> list[str]: @dataclass class FedEntry: term: str - layer: str # guild | channel | member + layer: str # guild | channel | member entity: str # channel_id (channel layer), user_id (member layer), "" (guild layer) definition: str synonyms: list[str] = field(default_factory=list) inferred: bool = False + kind: str = "" # metric | table | rule | dimension + applies_to: str = "" # 관련 테이블/컬럼 (예: users, orders.amount) + tags: list[str] = field(default_factory=list) def __post_init__(self) -> None: if not isinstance(self.synonyms, list): self.synonyms = _parse_synonyms(self.synonyms) + if not isinstance(self.tags, list): + self.tags = [t.strip() for t in str(self.tags).split(",") if t.strip()] def to_json(self) -> str: return json.dumps( { - "term": self.term, "layer": self.layer, "entity": self.entity, - "definition": self.definition, "synonyms": self.synonyms, + "term": self.term, + "layer": self.layer, + "entity": self.entity, + "definition": self.definition, + "synonyms": self.synonyms, "inferred": self.inferred, + "kind": self.kind, + "applies_to": self.applies_to, + "tags": self.tags, }, ensure_ascii=False, ) @@ -80,9 +94,15 @@ def to_json(self) -> str: def from_json(raw: str) -> "FedEntry": d = json.loads(raw) return FedEntry( - term=d["term"], layer=d["layer"], entity=d.get("entity", ""), - definition=d["definition"], synonyms=d.get("synonyms", []), + term=d["term"], + layer=d["layer"], + entity=d.get("entity", ""), + definition=d["definition"], + synonyms=d.get("synonyms", []), inferred=d.get("inferred", False), + kind=d.get("kind", ""), + applies_to=d.get("applies_to", ""), + tags=d.get("tags", []), ) @@ -117,6 +137,19 @@ def spec(self) -> ToolSpec: "type": "string", "description": "쉼표 구분 동의어 (예: active_user,활성화고객)", }, + "kind": { + "type": "string", + "enum": ["metric", "table", "rule", "dimension"], + "description": "용어 종류. metric=지표, table=테이블/엔티티, rule=비즈니스 규칙, dimension=분류 기준.", + }, + "applies_to": { + "type": "string", + "description": "관련 테이블 또는 컬럼 (예: users, orders.amount).", + }, + "tags": { + "type": "string", + "description": "쉼표 구분 태그 (예: growth,retention).", + }, "inferred": { "type": "boolean", "description": "true 시 LLM 추론 임시 정의로 표시. 사용자 확인 후 재등록 권장.", @@ -146,21 +179,32 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: channel_id = ctx.identity.effective_channel_id if args.get("list"): - return ToolResult(call_id="", content=_render_effective(ctx.store, scope, channel_id, user_id)) + return ToolResult( + call_id="", + content=_render_effective(ctx.store, scope, channel_id, user_id), + ) if args.get("scan"): return ToolResult(call_id="", content=_scan_schema(ctx.store, scope)) term = str(args.get("term", "")).strip() if not term: - return ToolResult(call_id="", content="❌ term 파라미터가 필요합니다.", is_error=True) + return ToolResult( + call_id="", content="❌ term 파라미터가 필요합니다.", is_error=True + ) if ":" in term: - return ToolResult(call_id="", content="❌ term에 ':'를 사용할 수 없습니다.", is_error=True) + return ToolResult( + call_id="", content="❌ term에 ':'를 사용할 수 없습니다.", is_error=True + ) if args.get("remove"): # 존재하는 항목 모두 삭제 — guild layer는 admin만 삭제 가능 deleted_tags: list[str] = [] - for lyr, ent in [("guild", ""), ("channel", channel_id), ("member", user_id)]: + for lyr, ent in [ + ("guild", ""), + ("channel", channel_id), + ("member", user_id), + ]: if lyr == "guild" and not ctx.identity.is_admin: continue k = _kv_key(term, lyr, ent) @@ -176,13 +220,21 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: content=f"⚠️ **{term}** — 전사(guild) 항목이 존재하지만 관리자만 삭제할 수 있습니다.", is_error=True, ) - return ToolResult(call_id="", content=f"⚠️ **{term}** — 등록된 정의가 없습니다.") + return ToolResult( + call_id="", content=f"⚠️ **{term}** — 등록된 정의가 없습니다." + ) if ctx.audit is not None: await ctx.audit.record( - AuditEvent(actor=user_id, action="term_custom_remove", - scope=scope, detail={"term": term, "layers": deleted_tags}) + AuditEvent( + actor=user_id, + action="term_custom_remove", + scope=scope, + detail={"term": term, "layers": deleted_tags}, + ) ) - return ToolResult(call_id="", content=f"🗑️ **{term}** [{', '.join(deleted_tags)}] 삭제") + return ToolResult( + call_id="", content=f"🗑️ **{term}** [{', '.join(deleted_tags)}] 삭제" + ) layer = str(args.get("layer", "member")).strip().lower() if layer not in _LAYERS: @@ -206,23 +258,45 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: is_error=True, ) - entity = "" if layer == "guild" else (user_id if layer == "member" else channel_id) + entity = ( + "" if layer == "guild" else (user_id if layer == "member" else channel_id) + ) key = _kv_key(term, layer, entity) definition = str(args.get("definition", "")).strip() if not definition: - return ToolResult(call_id="", content="❌ definition 파라미터가 필요합니다.", is_error=True) + return ToolResult( + call_id="", + content="❌ definition 파라미터가 필요합니다.", + is_error=True, + ) synonyms = _parse_synonyms(args.get("synonyms")) inferred = bool(args.get("inferred", False)) - - entry = FedEntry(term=term, layer=layer, entity=entity, - definition=definition, synonyms=synonyms, inferred=inferred) + kind = str(args.get("kind", "")).strip().lower() + applies_to = str(args.get("applies_to", "")).strip() + tags = [t.strip() for t in str(args.get("tags", "")).split(",") if t.strip()] + + entry = FedEntry( + term=term, + layer=layer, + entity=entity, + definition=definition, + synonyms=synonyms, + inferred=inferred, + kind=kind, + applies_to=applies_to, + tags=tags, + ) ctx.store.kv_set(scope, key, entry.to_json()) if ctx.audit is not None: await ctx.audit.record( - AuditEvent(actor=user_id, action="term_custom", - scope=scope, detail={"term": term, "layer": layer}) + AuditEvent( + actor=user_id, + action="term_custom", + scope=scope, + detail={"term": term, "layer": layer}, + ) ) tag = _layer_tag(layer, entity, user_id, channel_id) @@ -238,6 +312,7 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: # Helpers # --------------------------------------------------------------------------- + def _layer_tag(layer: str, entity: str, user_id: str, channel_id: str) -> str: if layer == "guild": return "전사" @@ -250,6 +325,7 @@ def _layer_tag(layer: str, entity: str, user_id: str, channel_id: str) -> str: # Schema scan # --------------------------------------------------------------------------- + def _scan_schema(store: Any, scope: str) -> str: col_entries = store.kv_list_prefix(scope, _ENRICH_PREFIX + ":") if not col_entries: @@ -304,6 +380,7 @@ def _scan_schema(store: Any, scope: str) -> str: # System-prompt helpers # --------------------------------------------------------------------------- + def _load_all(store: Any, scope: str) -> dict[str, list[FedEntry]]: """KV에서 모든 cterm 엔트리를 {term_lower: [FedEntry]} 로 반환.""" raw = store.kv_list_prefix(scope, _KV_PREFIX + ":") @@ -324,63 +401,91 @@ def _load_all(store: Any, scope: str) -> dict[str, list[FedEntry]]: return by_term -def build_prompt_section(store: Any, scope: str, channel_id: str, user_id: str) -> str: - """현재 채널 기준 narrow→wide lookup 용어 섹션 + 모호 용어 지침 반환.""" - by_term = _load_all(store, scope) - - if not by_term: - return _AMBIGUOUS_TERM_POLICY - - lines: list[str] = [] - for term_lower in sorted(by_term): - line = _resolve_term(by_term[term_lower], channel_id, user_id) - if line: - lines.append(line) - - header = ( - "## Business Terminology\n" - "(lookup 우선순위: 개인 > 채널(팀) > 전사)\n" - ) - body = "\n".join(lines) if lines else "(없음)" - return header + body + "\n\n" + _AMBIGUOUS_TERM_POLICY - +_KIND_SQL_HINT: dict[str, str] = { + "metric": "집계 지표 — SELECT/HAVING 절의 집계식으로 사용", + "rule": "비즈니스 규칙 — WHERE 절에 AND 조건으로 추가", + "dimension": "분류 기준 — GROUP BY 또는 SELECT 컬럼으로 활용", + "table": "테이블/엔티티 — FROM/JOIN 대상 선택 시 참고", +} +_KIND_ORDER = ["metric", "dimension", "rule", "table"] _AMBIGUOUS_TERM_POLICY = """\ ## Ambiguous Term Policy -사전에 없는 주관적/모호한 표현(예: 활성화고객, 신규고객, 우량고객)을 발견하면: -1. 현재 DB 스키마 컨텍스트에서 가장 합리적인 해석으로 SQL을 작성하고 실행한다. -2. 쿼리 후 사용한 해석을 명시하고, term_custom 등록 여부와 범위(guild/channel/member)를 사용자에게 묻는다. - 예: "'신규고객'을 'users.created_at >= NOW()-30일'로 해석했습니다. 이 정의를 어느 범위로 등록할까요?" -3. 사용자가 범위를 지정하면 term_custom 툴로 즉시 등록한다 (inferred=true). +사전에 없는 주관적/모호한 표현을 발견하면: +1. DB 스키마 기준으로 가장 합리적인 해석으로 SQL을 실행한다. +2. 실행 후 사용한 해석을 명시하고, kind(metric/rule/dimension/table)와 범위(guild/channel/member)를 사용자에게 묻는다. + 예: "'신규고객'을 'users.created_at >= NOW()-30일'로 해석했습니다. metric/rule/dimension/table 중 어느 종류이며, 어느 범위로 등록할까요?" +3. 사용자가 지정하면 term_custom 툴로 즉시 등록한다 (inferred=true). 4. inferred=true 엔트리가 이미 있으면 해당 정의를 우선 사용하되, 사용자에게 확정 여부를 확인한다.\ """ +def _resolve_entry( + entries: list[FedEntry], channel_id: str, user_id: str +) -> FedEntry | None: + """narrow→wide lookup: member > channel > guild. 승리 FedEntry 반환.""" + for e in entries: + if e.layer == "member" and e.entity == user_id: + return e + for e in entries: + if e.layer == "channel" and e.entity == channel_id: + return e + for e in entries: + if e.layer == "guild": + return e + return None + + +def _tag_for(e: FedEntry) -> str: + if e.layer == "member": + return f"개인:{e.entity}" + if e.layer == "channel": + return "채널" + return "전사" + + def _fmt_entry(e: FedEntry, tag: str) -> str: syns = ", ".join(e.synonyms) syn_str = f" (= {syns})" if syns else "" inferred_badge = " 🤖" if e.inferred else "" - return f"- **{e.term}** [{tag}]{syn_str}{inferred_badge}: {e.definition}" + kind_badge = f" `{e.kind}`" if e.kind else "" + return ( + f"- **{e.term}**{kind_badge} [{tag}]{syn_str}{inferred_badge}: {e.definition}" + ) -def _resolve_term(entries: list[FedEntry], channel_id: str, user_id: str) -> str: - """narrow→wide lookup: member > channel > guild.""" - # 1. 개인 오버라이드 - for e in entries: - if e.layer == "member" and e.entity == user_id: - return _fmt_entry(e, f"개인:{user_id}") +def build_prompt_section(store: Any, scope: str, channel_id: str, user_id: str) -> str: + """kind별로 그룹화된 시멘틱 용어 섹션 + 모호 용어 지침 반환.""" + by_term = _load_all(store, scope) - # 2. 이 채널 정의 - for e in entries: - if e.layer == "channel" and e.entity == channel_id: - return _fmt_entry(e, "채널") + if not by_term: + return _AMBIGUOUS_TERM_POLICY - # 3. 전사 공통 - for e in entries: - if e.layer == "guild": - return _fmt_entry(e, "전사") + groups: dict[str, list[str]] = {} + for term_lower in sorted(by_term): + e = _resolve_entry(by_term[term_lower], channel_id, user_id) + if e is None: + continue + k = e.kind if e.kind in _KIND_SQL_HINT else "" + groups.setdefault(k, []).append(_fmt_entry(e, _tag_for(e))) + + if not groups: + return _AMBIGUOUS_TERM_POLICY + + parts: list[str] = [ + "## Business Terminology\n(lookup 우선순위: 개인 > 채널(팀) > 전사)\n" + ] + for kind in _KIND_ORDER + [""]: + lines = groups.get(kind, []) + if not lines: + continue + if kind: + parts.append(f"### {kind.capitalize()}s — {_KIND_SQL_HINT[kind]}") + else: + parts.append("### 기타") + parts.extend(lines) - return "" + return "\n".join(parts) + "\n\n" + _AMBIGUOUS_TERM_POLICY def _render_effective(store: Any, scope: str, channel_id: str, user_id: str) -> str: @@ -391,9 +496,9 @@ def _render_effective(store: Any, scope: str, channel_id: str, user_id: str) -> lines = ["**Business Terminology — 현재 채널 기준 유효 정의**\n"] for term_lower in sorted(by_term): - line = _resolve_term(by_term[term_lower], channel_id, user_id) - if line: - lines.append(line) + e = _resolve_entry(by_term[term_lower], channel_id, user_id) + if e: + lines.append(_fmt_entry(e, _tag_for(e))) if len(lines) == 1: lines.append("(이 채널에 적용되는 용어 정의가 없습니다)") diff --git a/tests/test_adapters.py b/tests/test_adapters.py index 33e3096..4fa48d7 100644 --- a/tests/test_adapters.py +++ b/tests/test_adapters.py @@ -21,8 +21,18 @@ def test_audit_record_then_query() -> None: store = SqliteStore() - asyncio.run(store.record(AuditEvent(actor="u1", action="run_sql", scope="s", detail={"q": "SELECT 1"}))) - asyncio.run(store.record(AuditEvent(actor="u1", action="define_metric", scope="s", detail={}))) + asyncio.run( + store.record( + AuditEvent( + actor="u1", action="run_sql", scope="s", detail={"q": "SELECT 1"} + ) + ) + ) + asyncio.run( + store.record( + AuditEvent(actor="u1", action="define_metric", scope="s", detail={}) + ) + ) asyncio.run(store.record(AuditEvent(actor="other", action="run_sql", scope="s"))) events = asyncio.run(store.query("u1")) @@ -36,17 +46,23 @@ def test_audit_record_then_query() -> None: def test_session_save_then_load_reconstructs_transcript() -> None: store = SqliteStore() - identity = Identity(user_id="u1", guild_id="g", channel_id="c", thread_id="t", is_admin=True) + identity = Identity( + user_id="u1", guild_id="g", channel_id="c", thread_id="t", is_admin=True + ) session = Session(identity=identity) session.add(Message(role=Role.USER, content="hi")) session.add( Message( role=Role.ASSISTANT, content="", - tool_calls=[ToolCall(id="call_1", name="run_sql", arguments={"q": "SELECT 1"})], + tool_calls=[ + ToolCall(id="call_1", name="run_sql", arguments={"q": "SELECT 1"}) + ], ) ) - session.add(Message(role=Role.TOOL, content="ok", tool_call_id="call_1", name="run_sql")) + session.add( + Message(role=Role.TOOL, content="ok", tool_call_id="call_1", name="run_sql") + ) key = identity.session_key() asyncio.run(store.save(key, session)) @@ -102,7 +118,9 @@ def test_postgres_explorer_satisfies_protocol() -> None: def test_postgres_explorer_execute() -> None: explorer = PostgresExplorer("postgresql://ignored") - order_rows = asyncio.run(explorer.execute("SELECT * FROM orders WHERE status='paid'")) + order_rows = asyncio.run( + explorer.execute("SELECT * FROM orders WHERE status='paid'") + ) assert order_rows and "amount" in order_rows[0] capped = asyncio.run(explorer.execute("select * from orders", limit=1)) diff --git a/tests/test_bench_demo.py b/tests/test_bench_demo.py index 0af6c9e..5e3819c 100644 --- a/tests/test_bench_demo.py +++ b/tests/test_bench_demo.py @@ -48,8 +48,12 @@ def test_demo_federation_resolves_distinct_definitions(): mkt = demo._marketing_identity() fin = demo._finance_identity() - demo._define_term(store, demo.GUILD, "active_user", "channel", demo.CH_MARKETING, "30d login") - demo._define_term(store, demo.GUILD, "active_user", "channel", demo.CH_FINANCE, "paid sub") + demo._define_term( + store, demo.GUILD, "active_user", "channel", demo.CH_MARKETING, "30d login" + ) + demo._define_term( + store, demo.GUILD, "active_user", "channel", demo.CH_FINANCE, "paid sub" + ) mkt_rendered = _render_effective(store, demo.GUILD, demo.CH_MARKETING, mkt.user_id) fin_rendered = _render_effective(store, demo.GUILD, demo.CH_FINANCE, fin.user_id) diff --git a/tests/test_confirm_ingest.py b/tests/test_confirm_ingest.py new file mode 100644 index 0000000..fa5d222 --- /dev/null +++ b/tests/test_confirm_ingest.py @@ -0,0 +1,272 @@ +"""confirm_ingest — pending 후보 등록 및 OkfBundle 연동 테스트.""" + +from __future__ import annotations + +import asyncio +import json +import tempfile +from pathlib import Path +from typing import Sequence + +from lang2sql.adapters.storage.sqlite_store import SqliteStore +from lang2sql.core.identity import Identity +from lang2sql.core.ports.ingestion import CandidateKind, SemanticCandidate +from lang2sql.core.types import Completion, Message, ToolSpec +from lang2sql.harness.context import HarnessContext +from lang2sql.harness.session import Session +from lang2sql.harness.tool_registry import ToolRegistry +from lang2sql.tools import build_default_tools +from lang2sql.tools.confirm_ingest import ConfirmIngest, _dict_to_candidate, _select +from lang2sql.tools.ingest_doc import PENDING_PREFIX, IngestDoc, _candidate_to_dict +from lang2sql.tools.semantic_federation import FedEntry, _kv_key + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +class _FakeLLM: + def __init__(self, content: str = "[]") -> None: + self._content = content + + async def complete( + self, messages: Sequence[Message], tools: Sequence[ToolSpec] = () + ) -> Completion: + return Completion(content=self._content, finish_reason="stop") + + +def _make_ctx(store: SqliteStore, okf_bundle_dir: str | None = None) -> HarnessContext: + identity = Identity(user_id="u1", guild_id="g1", channel_id="c1") + from lang2sql.ingestion import FileSource, IngestionPipeline, LLMExtractor + from lang2sql.memory import ( + InjectAllRecall, + InMemoryStore, + ManualExtractor, + MemoryService, + ) + + memory = MemoryService(InMemoryStore(), InjectAllRecall(), ManualExtractor()) + ingestion = IngestionPipeline() + source = FileSource() + extractor = LLMExtractor(_FakeLLM()) + tools = ToolRegistry( + build_default_tools( + memory=memory, ingestion=ingestion, source=source, extractor=extractor + ) + ) + return HarnessContext( + identity=identity, + llm=_FakeLLM(), + tools=tools, + session=Session(identity=identity), + store=store, + okf_bundle_dir=okf_bundle_dir, + ) + + +def _seed_pending( + store: SqliteStore, scope: str, ref: str, candidates: list[SemanticCandidate] +) -> None: + key = f"{PENDING_PREFIX}:{ref}" + store.kv_set(scope, key, json.dumps([_candidate_to_dict(c) for c in candidates])) + + +_SAMPLE = [ + SemanticCandidate( + CandidateKind.METRIC, + "monthly_revenue", + "SUM(orders.amount)", + applies_to="orders", + ), + SemanticCandidate(CandidateKind.RULE, "exclude_cancelled", "status != 'cancelled'"), + SemanticCandidate(CandidateKind.DIMENSION, "customer_tier", "users.tier"), +] + + +# --------------------------------------------------------------------------- +# 직렬화 단위 테스트 +# --------------------------------------------------------------------------- + + +def test_candidate_roundtrip() -> None: + for c in _SAMPLE: + d = _candidate_to_dict(c) + restored = _dict_to_candidate(d) + assert restored.kind == c.kind + assert restored.name == c.name + assert restored.definition == c.definition + + +def test_select_all() -> None: + assert _select(_SAMPLE, "all") == _SAMPLE + + +def test_select_indices() -> None: + result = _select(_SAMPLE, "1,3") + assert result is not None + assert [c.name for c in result] == ["monthly_revenue", "customer_tier"] + + +def test_select_out_of_range_returns_none() -> None: + assert _select(_SAMPLE, "9") is None + + +def test_select_invalid_string_returns_none() -> None: + assert _select(_SAMPLE, "foo") is None + + +# --------------------------------------------------------------------------- +# confirm_ingest 동작 테스트 +# --------------------------------------------------------------------------- + + +def test_confirm_all_saves_fed_entries() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + scope = ctx.identity.kv_scope + _seed_pending(store, scope, "defs.md", _SAMPLE) + + tool = ConfirmIngest() + result = asyncio.run( + tool.run({"ref": "defs.md", "accept": "all", "layer": "guild"}, ctx) + ) + + assert not result.is_error + assert "3 term(s)" in result.content + + for cand in _SAMPLE: + raw = store.kv_get(scope, _kv_key(cand.name, "guild", "")) + assert raw is not None + entry = FedEntry.from_json(raw) + assert entry.term == cand.name + assert entry.kind == cand.kind.value + assert entry.layer == "guild" + + +def test_confirm_by_index_saves_selected_only() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + scope = ctx.identity.kv_scope + _seed_pending(store, scope, "defs.md", _SAMPLE) + + result = asyncio.run( + ConfirmIngest().run({"ref": "defs.md", "accept": "2", "layer": "guild"}, ctx) + ) + + assert not result.is_error + assert "1 term(s)" in result.content + + assert store.kv_get(scope, _kv_key("exclude_cancelled", "guild", "")) is not None + assert store.kv_get(scope, _kv_key("monthly_revenue", "guild", "")) is None + + +def test_confirm_channel_layer_uses_channel_entity() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + scope = ctx.identity.kv_scope + _seed_pending(store, scope, "defs.md", [_SAMPLE[0]]) + + asyncio.run( + ConfirmIngest().run( + {"ref": "defs.md", "accept": "all", "layer": "channel"}, ctx + ) + ) + + ch_id = ctx.identity.effective_channel_id + raw = store.kv_get(scope, _kv_key("monthly_revenue", "channel", ch_id)) + assert raw is not None + entry = FedEntry.from_json(raw) + assert entry.entity == ch_id + + +def test_confirm_member_layer_uses_user_id() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + scope = ctx.identity.kv_scope + _seed_pending(store, scope, "defs.md", [_SAMPLE[0]]) + + asyncio.run( + ConfirmIngest().run({"ref": "defs.md", "accept": "all", "layer": "member"}, ctx) + ) + + raw = store.kv_get( + scope, _kv_key("monthly_revenue", "member", ctx.identity.user_id) + ) + assert raw is not None + + +def test_confirm_clears_pending_key_after_success() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + scope = ctx.identity.kv_scope + _seed_pending(store, scope, "defs.md", [_SAMPLE[0]]) + + asyncio.run( + ConfirmIngest().run({"ref": "defs.md", "accept": "all", "layer": "guild"}, ctx) + ) + + assert store.kv_get(scope, f"{PENDING_PREFIX}:defs.md") is None + + +def test_confirm_missing_ref_returns_error() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + result = asyncio.run(ConfirmIngest().run({"ref": "nonexistent.md"}, ctx)) + assert result.is_error + assert "ingest_doc" in result.content + + +def test_confirm_no_ref_arg_returns_error() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + result = asyncio.run(ConfirmIngest().run({}, ctx)) + assert result.is_error + + +# --------------------------------------------------------------------------- +# ingest_doc → confirm_ingest 전체 연동 테스트 +# --------------------------------------------------------------------------- + + +def test_ingest_doc_saves_pending_key() -> None: + store = SqliteStore() + ctx = _make_ctx(store) + scope = ctx.identity.kv_scope + + candidates = [SemanticCandidate(CandidateKind.METRIC, "active_user", "30d login")] + _seed_pending(store, scope, "test.md", candidates) + + raw = store.kv_get(scope, f"{PENDING_PREFIX}:test.md") + assert raw is not None + loaded = json.loads(raw) + assert loaded[0]["name"] == "active_user" + + +def test_confirm_with_okf_bundle_exports_files() -> None: + store = SqliteStore() + with tempfile.TemporaryDirectory() as bundle_dir: + ctx = _make_ctx(store, okf_bundle_dir=bundle_dir) + scope = ctx.identity.kv_scope + _seed_pending(store, scope, "defs.md", [_SAMPLE[0]]) + + result = asyncio.run( + ConfirmIngest().run( + {"ref": "defs.md", "accept": "all", "layer": "guild"}, ctx + ) + ) + assert not result.is_error + + md_files = list(Path(bundle_dir).rglob("*.md")) + assert len(md_files) >= 1 + assert any("monthly_revenue" in f.name for f in md_files) + + +def test_registered_tool_name_in_registry() -> None: + from lang2sql.tenancy.concierge import ContextConcierge + + concierge = ContextConcierge() + identity = Identity(user_id="u1", guild_id="g1", channel_id="c1") + ctx = asyncio.run(concierge.build_context(identity)) + names = {s.name for s in ctx.tools.specs()} + assert "confirm_ingest" in names + assert "ingest_doc" in names diff --git a/tests/test_db_adapters.py b/tests/test_db_adapters.py index 4c4b629..ab09f30 100644 --- a/tests/test_db_adapters.py +++ b/tests/test_db_adapters.py @@ -19,9 +19,9 @@ explorer_from_env, ) - # --- factory routing ------------------------------------------------------- + def test_factory_routes_d1(): exp = build_explorer("d1://acct123/db456") assert isinstance(exp, D1Explorer) @@ -61,13 +61,18 @@ def test_explorer_from_env(monkeypatch): # --- SQLAlchemy explorer against real SQLite ------------------------------- + def _seed_sqlite(path: str) -> None: from sqlalchemy import create_engine, text eng = create_engine(f"sqlite:///{path}") with eng.begin() as conn: - conn.execute(text("CREATE TABLE users (id INTEGER PRIMARY KEY, email TEXT NOT NULL)")) - conn.execute(text("INSERT INTO users (id, email) VALUES (1, 'a@x.com'), (2, 'b@x.com')")) + conn.execute( + text("CREATE TABLE users (id INTEGER PRIMARY KEY, email TEXT NOT NULL)") + ) + conn.execute( + text("INSERT INTO users (id, email) VALUES (1, 'a@x.com'), (2, 'b@x.com')") + ) def test_sqlalchemy_explorer_introspect_and_execute(tmp_path): @@ -92,6 +97,7 @@ def test_sqlalchemy_explorer_introspect_and_execute(tmp_path): # --- D1 explorer with mocked HTTP transport -------------------------------- + def _d1_transport(sql, params): """Fake the D1 HTTP API: shape responses by the SQL it receives.""" s = sql.lower() @@ -99,12 +105,30 @@ def _d1_transport(sql, params): results = [{"name": "orders"}, {"name": "users"}] elif "pragma table_info" in s: results = [ - {"cid": 0, "name": "id", "type": "INTEGER", "notnull": 1, "dflt_value": None, "pk": 1}, - {"cid": 1, "name": "amount", "type": "REAL", "notnull": 0, "dflt_value": None, "pk": 0}, + { + "cid": 0, + "name": "id", + "type": "INTEGER", + "notnull": 1, + "dflt_value": None, + "pk": 1, + }, + { + "cid": 1, + "name": "amount", + "type": "REAL", + "notnull": 0, + "dflt_value": None, + "pk": 0, + }, ] else: results = [{"id": 1, "amount": 9.5}] - return {"success": True, "result": [{"results": results, "success": True}], "errors": []} + return { + "success": True, + "result": [{"results": results, "success": True}], + "errors": [], + } def test_d1_list_describe_execute(): diff --git a/tests/test_discord.py b/tests/test_discord.py index 2ac606c..1436f93 100644 --- a/tests/test_discord.py +++ b/tests/test_discord.py @@ -26,7 +26,6 @@ from lang2sql.frontends.discord.render import MAX_INLINE_ROWS from lang2sql.tenancy.concierge import ContextConcierge - # -- session_router ------------------------------------------------------- @@ -112,7 +111,12 @@ def test_term_custom_then_list() -> None: ) async def scenario() -> tuple[str, str]: - defined = await handlers.term_custom(ident, term="active_user", definition="logged in within 30 days", layer="channel") + defined = await handlers.term_custom( + ident, + term="active_user", + definition="logged in within 30 days", + layer="channel", + ) shown = await handlers.term_custom(ident, list_all=True) return defined.text, shown.text @@ -124,7 +128,9 @@ async def scenario() -> tuple[str, str]: def test_term_custom_list_empty_scope() -> None: handlers = CommandHandlers(ContextConcierge()) - ident = to_identity(InteractionContext(user_id="solo", guild_id="g9", channel_id="c9")) + ident = to_identity( + InteractionContext(user_id="solo", guild_id="g9", channel_id="c9") + ) shown = asyncio.run(handlers.term_custom(ident, list_all=True)) assert shown.text # empty scope returns some message @@ -132,11 +138,17 @@ def test_term_custom_list_empty_scope() -> None: def test_term_custom_is_scope_isolated() -> None: """A channel definition must not leak into a different channel (federation).""" handlers = CommandHandlers(ContextConcierge()) - marketing = to_identity(InteractionContext(user_id="u1", guild_id="g1", channel_id="mkt")) - product = to_identity(InteractionContext(user_id="u1", guild_id="g1", channel_id="prd")) + marketing = to_identity( + InteractionContext(user_id="u1", guild_id="g1", channel_id="mkt") + ) + product = to_identity( + InteractionContext(user_id="u1", guild_id="g1", channel_id="prd") + ) async def scenario() -> str: - await handlers.term_custom(marketing, term="active_user", definition="30d login", layer="channel") + await handlers.term_custom( + marketing, term="active_user", definition="30d login", layer="channel" + ) return (await handlers.term_custom(product, list_all=True)).text assert "active_user" not in asyncio.run(scenario()) @@ -144,7 +156,9 @@ async def scenario() -> str: def test_remember_and_audit_me() -> None: handlers = CommandHandlers(ContextConcierge()) - ident = to_identity(InteractionContext(user_id="u2", guild_id="g1", channel_id="c1")) + ident = to_identity( + InteractionContext(user_id="u2", guild_id="g1", channel_id="c1") + ) async def scenario() -> tuple[str, str]: remembered = await handlers.remember(ident, "prefers ISO dates") @@ -158,7 +172,9 @@ async def scenario() -> tuple[str, str]: def test_audit_me_empty() -> None: handlers = CommandHandlers(ContextConcierge()) - ident = to_identity(InteractionContext(user_id="never-acted", guild_id="g1", channel_id="c1")) + ident = to_identity( + InteractionContext(user_id="never-acted", guild_id="g1", channel_id="c1") + ) audit = asyncio.run(handlers.audit_me(ident)) assert "No audited activity" in audit.text @@ -166,7 +182,9 @@ def test_audit_me_empty() -> None: def test_query_returns_outbound_message() -> None: """With the default FakeLLM (no OPENAI key), a query still returns text.""" handlers = CommandHandlers(ContextConcierge()) - ident = to_identity(InteractionContext(user_id="u3", guild_id="g1", channel_id="c1")) + ident = to_identity( + InteractionContext(user_id="u3", guild_id="g1", channel_id="c1") + ) out = asyncio.run(handlers.query(ident, "how many users signed up?")) assert isinstance(out.text, str) assert out.text # non-empty @@ -175,7 +193,9 @@ def test_query_returns_outbound_message() -> None: def test_query_persists_session() -> None: concierge = ContextConcierge() handlers = CommandHandlers(concierge) - ident = to_identity(InteractionContext(user_id="u4", guild_id="g1", channel_id="c1")) + ident = to_identity( + InteractionContext(user_id="u4", guild_id="g1", channel_id="c1") + ) async def scenario(): await handlers.query(ident, "first question") @@ -189,7 +209,9 @@ async def scenario(): def test_connect_stub_acknowledges() -> None: concierge = ContextConcierge() handlers = CommandHandlers(concierge) - ident = to_identity(InteractionContext(user_id="u5", guild_id="g1", channel_id="c1")) + ident = to_identity( + InteractionContext(user_id="u5", guild_id="g1", channel_id="c1") + ) out = asyncio.run(handlers.connect(ident, "postgresql://localhost/db")) assert "saved" in out.text.lower() assert concierge.store.kv_get("g1", "dsn") == "postgresql://localhost/db" @@ -197,7 +219,9 @@ def test_connect_stub_acknowledges() -> None: def test_ingest_lists_or_reports() -> None: handlers = CommandHandlers(ContextConcierge()) - ident = to_identity(InteractionContext(user_id="u6", guild_id="g1", channel_id="c1")) + ident = to_identity( + InteractionContext(user_id="u6", guild_id="g1", channel_id="c1") + ) out = asyncio.run( handlers.ingest(ident, content="total_revenue is the sum of order amounts") ) diff --git a/tests/test_edge_cases.py b/tests/test_edge_cases.py index c390cf0..72d3f73 100644 --- a/tests/test_edge_cases.py +++ b/tests/test_edge_cases.py @@ -9,16 +9,21 @@ from lang2sql.frontends.discord.render import OutboundMessage from lang2sql.tools.semantic_federation import FedEntry, _kv_key, _render_effective - # -- FedEntry synonyms: string stored in KV (pre-fix data or JSON null) ------- def test_fed_entry_from_json_coerces_string_synonyms() -> None: """If old KV data has synonyms as a JSON string, from_json must produce a list.""" - raw_json = json.dumps({ - "term": "active_user", "layer": "guild", "entity": "", - "definition": "30d login", "synonyms": "활성유저, active", "inferred": False, - }) + raw_json = json.dumps( + { + "term": "active_user", + "layer": "guild", + "entity": "", + "definition": "30d login", + "synonyms": "활성유저, active", + "inferred": False, + } + ) entry = FedEntry.from_json(raw_json) assert isinstance(entry.synonyms, list), "synonyms must be a list after from_json" assert "활성유저" in entry.synonyms @@ -27,10 +32,16 @@ def test_fed_entry_from_json_coerces_string_synonyms() -> None: def test_fed_entry_from_json_handles_null_synonyms() -> None: """If KV data has synonyms=null, from_json must produce an empty list.""" - raw_json = json.dumps({ - "term": "revenue", "layer": "guild", "entity": "", - "definition": "gross revenue", "synonyms": None, "inferred": False, - }) + raw_json = json.dumps( + { + "term": "revenue", + "layer": "guild", + "entity": "", + "definition": "gross revenue", + "synonyms": None, + "inferred": False, + } + ) entry = FedEntry.from_json(raw_json) assert entry.synonyms == [] @@ -40,10 +51,16 @@ def test_render_effective_string_synonyms_in_kv_does_not_character_join() -> Non store = SqliteStore() scope = "g1" # Simulate a KV entry written by old code (synonyms as JSON string) - bad_json = json.dumps({ - "term": "active_user", "layer": "guild", "entity": "", - "definition": "30d login", "synonyms": "활성유저, active", "inferred": False, - }) + bad_json = json.dumps( + { + "term": "active_user", + "layer": "guild", + "entity": "", + "definition": "30d login", + "synonyms": "활성유저, active", + "inferred": False, + } + ) store.kv_set(scope, _kv_key("active_user", "guild", ""), bad_json) rendered = _render_effective(store, scope, "", "u1") assert "active_user" in rendered @@ -88,16 +105,20 @@ def test_term_custom_remove_emits_audit_event() -> None: ctx = asyncio.run(concierge.build_context(ident)) # Write then remove - asyncio.run(SemanticFederationTool().run( - {"term": "active_user", "definition": "30d login", "layer": "guild"}, ctx - )) - asyncio.run(SemanticFederationTool().run( - {"term": "active_user", "remove": True}, ctx - )) + asyncio.run( + SemanticFederationTool().run( + {"term": "active_user", "definition": "30d login", "layer": "guild"}, ctx + ) + ) + asyncio.run( + SemanticFederationTool().run({"term": "active_user", "remove": True}, ctx) + ) events = asyncio.run(ctx.audit.query(ident.user_id)) - assert any(e.action == "term_custom_remove" and e.detail.get("term") == "active_user" - for e in events) + assert any( + e.action == "term_custom_remove" and e.detail.get("term") == "active_user" + for e in events + ) # -- guild layer admin guard --------------------------------------------------- @@ -113,9 +134,11 @@ def test_guild_write_requires_admin() -> None: ident = Identity(user_id="u1", guild_id="g1", channel_id="c1", is_admin=False) ctx = asyncio.run(concierge.build_context(ident)) - result = asyncio.run(SemanticFederationTool().run( - {"term": "revenue", "definition": "gross revenue", "layer": "guild"}, ctx - )) + result = asyncio.run( + SemanticFederationTool().run( + {"term": "revenue", "definition": "gross revenue", "layer": "guild"}, ctx + ) + ) assert result.is_error assert "관리자" in result.content @@ -129,23 +152,35 @@ def test_guild_remove_non_admin_skips_guild_keeps_own_entry() -> None: concierge = ContextConcierge() # Admin registers the guild-layer term - admin_ctx = asyncio.run(concierge.build_context( - Identity(user_id="admin", guild_id="g1", channel_id="c1", is_admin=True) - )) - asyncio.run(SemanticFederationTool().run( - {"term": "revenue", "definition": "gross revenue", "layer": "guild"}, admin_ctx - )) + admin_ctx = asyncio.run( + concierge.build_context( + Identity(user_id="admin", guild_id="g1", channel_id="c1", is_admin=True) + ) + ) + asyncio.run( + SemanticFederationTool().run( + {"term": "revenue", "definition": "gross revenue", "layer": "guild"}, + admin_ctx, + ) + ) # Non-admin adds their own member-layer override - member_ctx = asyncio.run(concierge.build_context( - Identity(user_id="u1", guild_id="g1", channel_id="c1", is_admin=False) - )) - asyncio.run(SemanticFederationTool().run( - {"term": "revenue", "definition": "my override", "layer": "member"}, member_ctx - )) + member_ctx = asyncio.run( + concierge.build_context( + Identity(user_id="u1", guild_id="g1", channel_id="c1", is_admin=False) + ) + ) + asyncio.run( + SemanticFederationTool().run( + {"term": "revenue", "definition": "my override", "layer": "member"}, + member_ctx, + ) + ) # Non-admin removes — must keep guild entry, delete own member entry - asyncio.run(SemanticFederationTool().run({"term": "revenue", "remove": True}, member_ctx)) + asyncio.run( + SemanticFederationTool().run({"term": "revenue", "remove": True}, member_ctx) + ) scope = "g1" assert member_ctx.store.kv_get(scope, _kv_key("revenue", "guild", "")) is not None @@ -168,18 +203,30 @@ def test_channel_layer_term_visible_from_thread_context() -> None: from lang2sql.core.identity import Identity channel_ident = Identity(user_id="u1", guild_id="g1", channel_id="c1") - thread_ident = Identity(user_id="u2", guild_id="g1", channel_id="c1", thread_id="t1") + thread_ident = Identity( + user_id="u2", guild_id="g1", channel_id="c1", thread_id="t1" + ) # Both identities must resolve to the same channel entity - assert channel_ident.effective_channel_id == thread_ident.effective_channel_id == "c1" + assert ( + channel_ident.effective_channel_id == thread_ident.effective_channel_id == "c1" + ) store = SqliteStore() scope = "g1" - store.kv_set(scope, _kv_key("active_user", "channel", "c1"), - FedEntry(term="active_user", layer="channel", entity="c1", - definition="30d login").to_json()) + store.kv_set( + scope, + _kv_key("active_user", "channel", "c1"), + FedEntry( + term="active_user", layer="channel", entity="c1", definition="30d login" + ).to_json(), + ) # Term visible from channel context - assert "active_user" in _render_effective(store, scope, channel_ident.effective_channel_id, "u1") + assert "active_user" in _render_effective( + store, scope, channel_ident.effective_channel_id, "u1" + ) # Term also visible from thread context (inherits parent channel) - assert "active_user" in _render_effective(store, scope, thread_ident.effective_channel_id, "u2") + assert "active_user" in _render_effective( + store, scope, thread_ident.effective_channel_id, "u2" + ) diff --git a/tests/test_integration.py b/tests/test_integration.py index 3556cb5..79bb110 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -26,7 +26,17 @@ def _ctx(): def test_v1_tools_registered(): _, ctx = _ctx() names = {s.name for s in ctx.tools.specs()} - assert names == {"run_sql", "explore_schema", "enrich_schema", "term_custom", "org_setup", "ask_user", "remember", "ingest_doc"} + assert names == { + "run_sql", + "explore_schema", + "enrich_schema", + "term_custom", + "org_setup", + "ask_user", + "remember", + "ingest_doc", + "confirm_ingest", + } def test_run_sql_passes_gate_and_returns_rows(): @@ -50,11 +60,16 @@ def test_run_sql_tolerates_bad_limit(): def test_term_custom_is_scope_local(): from lang2sql.tools.semantic_federation import _render_effective + ident, ctx = _ctx() - asyncio.run(SemanticFederationTool().run( - {"term": "active_user", "definition": "30d login", "layer": "channel"}, ctx - )) - rendered = _render_effective(ctx.store, ident.kv_scope, ident.effective_channel_id, ident.user_id) + asyncio.run( + SemanticFederationTool().run( + {"term": "active_user", "definition": "30d login", "layer": "channel"}, ctx + ) + ) + rendered = _render_effective( + ctx.store, ident.kv_scope, ident.effective_channel_id, ident.user_id + ) assert "active_user" in rendered # a different channel does not see this channel-level definition other_rendered = _render_effective(ctx.store, ident.kv_scope, "c-fin", "u1") @@ -65,11 +80,15 @@ def test_term_custom_emits_audit_event(): concierge = ContextConcierge() ident = Identity(user_id="u1", guild_id="g1", channel_id="c-mkt", is_admin=True) ctx = asyncio.run(concierge.build_context(ident)) - asyncio.run(SemanticFederationTool().run( - {"term": "revenue", "definition": "gross revenue", "layer": "guild"}, ctx - )) + asyncio.run( + SemanticFederationTool().run( + {"term": "revenue", "definition": "gross revenue", "layer": "guild"}, ctx + ) + ) events = asyncio.run(ctx.audit.query(ident.user_id)) - assert any(e.action == "term_custom" and e.detail.get("term") == "revenue" for e in events) + assert any( + e.action == "term_custom" and e.detail.get("term") == "revenue" for e in events + ) def test_safety_pipeline_on_context(): diff --git a/tests/test_okf_bundle.py b/tests/test_okf_bundle.py new file mode 100644 index 0000000..99c2ec7 --- /dev/null +++ b/tests/test_okf_bundle.py @@ -0,0 +1,220 @@ +"""OkfBundle — export/import round-trip 및 파일 구조 테스트.""" + +from __future__ import annotations + +import tempfile +from pathlib import Path + +from lang2sql.adapters.storage.okf_bundle import OkfBundle, _entry_to_md, _md_to_entry +from lang2sql.adapters.storage.sqlite_store import SqliteStore +from lang2sql.tools.semantic_federation import FedEntry, _kv_key + + +def _populate(store: SqliteStore, scope: str, entries: list[FedEntry]) -> None: + for e in entries: + store.kv_set(scope, _kv_key(e.term, e.layer, e.entity), e.to_json()) + + +# ------------------------------------------------------------------ +# 직렬화 단위 테스트 +# ------------------------------------------------------------------ + + +def test_entry_to_md_contains_required_okf_fields() -> None: + entry = FedEntry( + term="활성고객", + layer="guild", + entity="", + definition="30일 내 로그인한 users", + kind="metric", + applies_to="users", + tags=["growth"], + ) + md = _entry_to_md(entry) + assert "type: Metric" in md + assert "title: 활성고객" in md + assert "description:" in md + assert "layer: guild" in md + + +def test_md_to_entry_roundtrip() -> None: + entry = FedEntry( + term="순매출", + layer="channel", + entity="mkt-123", + definition="환불 제외 매출", + synonyms=["net revenue"], + kind="metric", + applies_to="orders", + tags=["finance"], + inferred=True, + ) + with tempfile.NamedTemporaryFile( + suffix=".md", mode="w", delete=False, encoding="utf-8" + ) as f: + f.write(_entry_to_md(entry)) + tmp = Path(f.name) + + restored = _md_to_entry(tmp) + assert restored is not None + assert restored.term == "순매출" + assert restored.kind == "metric" + assert restored.layer == "channel" + assert restored.entity == "mkt-123" + assert restored.applies_to == "orders" + assert restored.tags == ["finance"] + assert restored.inferred is True + tmp.unlink() + + +def test_md_to_entry_unknown_type_becomes_empty_kind() -> None: + md = "---\ntype: Playbook\ntitle: foo\ndescription: bar\nlayer: guild\nentity: ''\ninferred: false\n---\n\nbar\n" + with tempfile.NamedTemporaryFile( + suffix=".md", mode="w", delete=False, encoding="utf-8" + ) as f: + f.write(md) + tmp = Path(f.name) + entry = _md_to_entry(tmp) + assert entry is not None + assert entry.kind == "" + tmp.unlink() + + +def test_md_to_entry_no_frontmatter_returns_none() -> None: + with tempfile.NamedTemporaryFile( + suffix=".md", mode="w", delete=False, encoding="utf-8" + ) as f: + f.write("no frontmatter here") + tmp = Path(f.name) + assert _md_to_entry(tmp) is None + tmp.unlink() + + +# ------------------------------------------------------------------ +# export / import 통합 테스트 +# ------------------------------------------------------------------ + + +def test_export_creates_kind_based_folders() -> None: + store = SqliteStore() + scope = "g1" + entries = [ + FedEntry("활성고객", "guild", "", "30일 로그인", kind="metric"), + FedEntry("orders", "guild", "", "주문 테이블", kind="table"), + FedEntry("환불제외", "guild", "", "status != refunded", kind="rule"), + FedEntry("고객등급", "guild", "", "users.tier", kind="dimension"), + FedEntry("기타용어", "guild", "", "정의 없음", kind=""), + ] + _populate(store, scope, entries) + + with tempfile.TemporaryDirectory() as tmp: + bundle = OkfBundle(tmp) + count = bundle.export(store, scope) + + assert count == 5 + assert (Path(tmp) / "guild" / "metrics" / "활성고객.md").exists() + assert (Path(tmp) / "guild" / "tables" / "orders.md").exists() + assert (Path(tmp) / "guild" / "rules" / "환불제외.md").exists() + assert (Path(tmp) / "guild" / "dimensions" / "고객등급.md").exists() + assert (Path(tmp) / "guild" / "misc" / "기타용어.md").exists() + + +def test_export_separates_scopes() -> None: + store = SqliteStore() + scope = "g1" + _populate( + store, + scope, + [ + FedEntry("활성고객", "guild", "", "30일 로그인", kind="metric"), + FedEntry("활성고객", "channel", "mkt", "7일 구매", kind="metric"), + ], + ) + + with tempfile.TemporaryDirectory() as tmp: + bundle = OkfBundle(tmp) + bundle.export(store, scope) + + assert (Path(tmp) / "guild" / "metrics" / "활성고객.md").exists() + assert (Path(tmp) / "channel:mkt" / "metrics" / "활성고객.md").exists() + + +def test_import_restores_kv_from_files() -> None: + store = SqliteStore() + scope = "g1" + original = FedEntry( + "순매출", + "guild", + "", + "환불 제외 매출", + kind="metric", + applies_to="orders", + tags=["finance"], + ) + _populate(store, scope, [original]) + + with tempfile.TemporaryDirectory() as tmp: + bundle = OkfBundle(tmp) + bundle.export(store, scope) + + # KV 비우고 import + empty_store = SqliteStore() + count = bundle.import_(empty_store, scope) + + assert count == 1 + key = _kv_key("순매출", "guild", "") + raw = empty_store.kv_get(scope, key) + assert raw is not None + restored = FedEntry.from_json(raw) + assert restored.term == "순매출" + assert restored.kind == "metric" + assert restored.applies_to == "orders" + + +def test_import_skips_reserved_files() -> None: + with tempfile.TemporaryDirectory() as tmp: + guild_dir = Path(tmp) / "guild" + guild_dir.mkdir() + (guild_dir / "index.md").write_text("# index", encoding="utf-8") + (guild_dir / "log.md").write_text("# log", encoding="utf-8") + + store = SqliteStore() + bundle = OkfBundle(tmp) + count = bundle.import_(store, "g1") + assert count == 0 + + +def test_full_roundtrip_preserves_all_fields() -> None: + store = SqliteStore() + scope = "g1" + original = FedEntry( + term="월매출", + layer="member", + entity="user-99", + definition="당월 발생 매출 합계", + synonyms=["monthly revenue"], + inferred=False, + kind="metric", + applies_to="orders.amount", + tags=["finance", "monthly"], + ) + _populate(store, scope, [original]) + + with tempfile.TemporaryDirectory() as tmp: + bundle = OkfBundle(tmp) + bundle.export(store, scope) + restored_store = SqliteStore() + bundle.import_(restored_store, scope) + + key = _kv_key("월매출", "member", "user-99") + raw = restored_store.kv_get(scope, key) + assert raw is not None + restored = FedEntry.from_json(raw) + + assert restored.term == "월매출" + assert restored.layer == "member" + assert restored.entity == "user-99" + assert restored.kind == "metric" + assert restored.applies_to == "orders.amount" + assert set(restored.tags) == {"finance", "monthly"} + assert restored.synonyms == ["monthly revenue"] diff --git a/tests/test_persistence.py b/tests/test_persistence.py index 2b1f559..f101187 100644 --- a/tests/test_persistence.py +++ b/tests/test_persistence.py @@ -24,7 +24,9 @@ def test_kv_federation_survives_new_instance(tmp_path) -> None: scope = "g1" writer = SqliteStore(db) - entry = FedEntry(term="revenue", layer="guild", entity="", definition="sum of order totals") + entry = FedEntry( + term="revenue", layer="guild", entity="", definition="sum of order totals" + ) writer.kv_set(scope, _kv_key("revenue", "guild", ""), entry.to_json()) writer.close() @@ -40,8 +42,16 @@ def test_kv_channel_overrides_guild_persisted(tmp_path) -> None: scope = "g1" store = SqliteStore(db) - store.kv_set(scope, _kv_key("active_user", "guild", ""), FedEntry("active_user", "guild", "", "guild def").to_json()) - store.kv_set(scope, _kv_key("active_user", "channel", "c1"), FedEntry("active_user", "channel", "c1", "channel def").to_json()) + store.kv_set( + scope, + _kv_key("active_user", "guild", ""), + FedEntry("active_user", "guild", "", "guild def").to_json(), + ) + store.kv_set( + scope, + _kv_key("active_user", "channel", "c1"), + FedEntry("active_user", "channel", "c1", "channel def").to_json(), + ) store.close() reader = SqliteStore(db) @@ -64,7 +74,9 @@ def test_encrypted_secrets_round_trip_and_ciphertext(tmp_path) -> None: assert blob is not None assert "postgresql" not in blob assert blob != "postgresql://u:p@host/db" - assert Fernet(key).decrypt(blob.encode("ascii")).decode() == "postgresql://u:p@host/db" + assert ( + Fernet(key).decrypt(blob.encode("ascii")).decode() == "postgresql://u:p@host/db" + ) asyncio.run(secrets.delete("guild:1", "dsn")) assert asyncio.run(secrets.get("guild:1", "dsn")) is None diff --git a/tests/test_semantic.py b/tests/test_semantic.py index f24c127..f6e7a1f 100644 --- a/tests/test_semantic.py +++ b/tests/test_semantic.py @@ -30,8 +30,16 @@ def _store_with_entries(entries: list[tuple[str, str, str, str]]) -> SqliteStore def test_channel_overrides_guild() -> None: store = SqliteStore() scope = "g1" - store.kv_set(scope, _kv_key("active_user", "guild", ""), FedEntry("active_user", "guild", "", "30d login").to_json()) - store.kv_set(scope, _kv_key("active_user", "channel", "c1"), FedEntry("active_user", "channel", "c1", "7d core action").to_json()) + store.kv_set( + scope, + _kv_key("active_user", "guild", ""), + FedEntry("active_user", "guild", "", "30d login").to_json(), + ) + store.kv_set( + scope, + _kv_key("active_user", "channel", "c1"), + FedEntry("active_user", "channel", "c1", "7d core action").to_json(), + ) rendered = _render_effective(store, scope, "c1", "u1") assert "7d core action" in rendered @@ -41,7 +49,11 @@ def test_channel_overrides_guild() -> None: def test_guild_fills_gap_when_channel_missing() -> None: store = SqliteStore() scope = "g1" - store.kv_set(scope, _kv_key("revenue", "guild", ""), FedEntry("revenue", "guild", "", "net revenue").to_json()) + store.kv_set( + scope, + _kv_key("revenue", "guild", ""), + FedEntry("revenue", "guild", "", "net revenue").to_json(), + ) rendered = _render_effective(store, scope, "c1", "u1") assert "net revenue" in rendered @@ -50,9 +62,21 @@ def test_guild_fills_gap_when_channel_missing() -> None: def test_member_overrides_channel_and_guild() -> None: store = SqliteStore() scope = "g1" - store.kv_set(scope, _kv_key("active_user", "guild", ""), FedEntry("active_user", "guild", "", "guild def").to_json()) - store.kv_set(scope, _kv_key("active_user", "channel", "c1"), FedEntry("active_user", "channel", "c1", "channel def").to_json()) - store.kv_set(scope, _kv_key("active_user", "member", "u1"), FedEntry("active_user", "member", "u1", "member def").to_json()) + store.kv_set( + scope, + _kv_key("active_user", "guild", ""), + FedEntry("active_user", "guild", "", "guild def").to_json(), + ) + store.kv_set( + scope, + _kv_key("active_user", "channel", "c1"), + FedEntry("active_user", "channel", "c1", "channel def").to_json(), + ) + store.kv_set( + scope, + _kv_key("active_user", "member", "u1"), + FedEntry("active_user", "member", "u1", "member def").to_json(), + ) rendered = _render_effective(store, scope, "c1", "u1") assert "member def" in rendered @@ -63,8 +87,16 @@ def test_member_overrides_channel_and_guild() -> None: def test_two_channels_isolated() -> None: store = SqliteStore() scope = "g1" - store.kv_set(scope, _kv_key("active_user", "channel", "mkt"), FedEntry("active_user", "channel", "mkt", "30d login").to_json()) - store.kv_set(scope, _kv_key("active_user", "channel", "fin"), FedEntry("active_user", "channel", "fin", "paid subscriber").to_json()) + store.kv_set( + scope, + _kv_key("active_user", "channel", "mkt"), + FedEntry("active_user", "channel", "mkt", "30d login").to_json(), + ) + store.kv_set( + scope, + _kv_key("active_user", "channel", "fin"), + FedEntry("active_user", "channel", "fin", "paid subscriber").to_json(), + ) mkt = _render_effective(store, scope, "mkt", "u1") fin = _render_effective(store, scope, "fin", "u2") @@ -84,3 +116,157 @@ def test_build_prompt_section_includes_ambiguous_term_policy() -> None: store = SqliteStore() section = build_prompt_section(store, "g1", "c1", "u1") assert "Ambiguous Term Policy" in section + + +def test_fed_entry_kind_applies_to_tags_roundtrip() -> None: + entry = FedEntry( + term="활성고객", + layer="guild", + entity="", + definition="30일 내 로그인한 users", + kind="metric", + applies_to="users", + tags=["growth", "retention"], + ) + restored = FedEntry.from_json(entry.to_json()) + assert restored.kind == "metric" + assert restored.applies_to == "users" + assert restored.tags == ["growth", "retention"] + + +def test_fed_entry_backward_compat_missing_new_fields() -> None: + # kind/applies_to/tags 없는 기존 JSON도 파싱 가능해야 함 + import json + + old_json = json.dumps( + { + "term": "revenue", + "layer": "guild", + "entity": "", + "definition": "net revenue", + "synonyms": [], + "inferred": False, + } + ) + entry = FedEntry.from_json(old_json) + assert entry.kind == "" + assert entry.applies_to == "" + assert entry.tags == [] + + +def test_fmt_entry_shows_kind_badge() -> None: + from lang2sql.tools.semantic_federation import _fmt_entry + + entry = FedEntry( + term="활성고객", + layer="guild", + entity="", + definition="30일 내 로그인", + kind="metric", + ) + rendered = _fmt_entry(entry, "전사") + assert "`metric`" in rendered + + +# --------------------------------------------------------------------------- +# PR4: kind-grouped prompt section + disambiguation policy +# --------------------------------------------------------------------------- + + +def _seed(store: SqliteStore, scope: str, entries: list[FedEntry]) -> None: + for e in entries: + store.kv_set(scope, _kv_key(e.term, e.layer, e.entity), e.to_json()) + + +def test_prompt_section_groups_by_kind() -> None: + store = SqliteStore() + _seed( + store, + "g1", + [ + FedEntry("월매출", "guild", "", "SUM(orders.amount)", kind="metric"), + FedEntry("환불제외", "guild", "", "status != refunded", kind="rule"), + FedEntry("고객등급", "guild", "", "users.tier", kind="dimension"), + ], + ) + section = build_prompt_section(store, "g1", "c1", "u1") + + assert "### Metrics" in section + assert "### Rules" in section + assert "### Dimensions" in section + assert "월매출" in section + assert "환불제외" in section + assert "고객등급" in section + + +def test_prompt_section_kind_headers_contain_sql_hint() -> None: + store = SqliteStore() + _seed( + store, + "g1", + [FedEntry("월매출", "guild", "", "SUM(orders.amount)", kind="metric")], + ) + section = build_prompt_section(store, "g1", "c1", "u1") + + assert "SELECT/HAVING" in section + assert "### Metrics" in section + + +def test_prompt_section_unknown_kind_goes_to_기타() -> None: + store = SqliteStore() + _seed(store, "g1", [FedEntry("알수없음", "guild", "", "정의 없음", kind="")]) + section = build_prompt_section(store, "g1", "c1", "u1") + + assert "### 기타" in section + assert "알수없음" in section + + +def test_prompt_section_skips_empty_kind_groups() -> None: + store = SqliteStore() + _seed( + store, + "g1", + [FedEntry("월매출", "guild", "", "SUM(orders.amount)", kind="metric")], + ) + section = build_prompt_section(store, "g1", "c1", "u1") + + assert "### Rules" not in section + assert "### Dimensions" not in section + + +def test_ambiguous_policy_mentions_kind() -> None: + store = SqliteStore() + section = build_prompt_section(store, "g1", "c1", "u1") + assert "metric/rule/dimension/table" in section + + +def test_resolve_entry_member_wins_over_channel() -> None: + from lang2sql.tools.semantic_federation import _resolve_entry + + entries = [ + FedEntry("t", "guild", "", "guild-def"), + FedEntry("t", "channel", "c1", "channel-def"), + FedEntry("t", "member", "u1", "member-def"), + ] + result = _resolve_entry(entries, "c1", "u1") + assert result is not None + assert result.definition == "member-def" + + +def test_resolve_entry_channel_wins_over_guild() -> None: + from lang2sql.tools.semantic_federation import _resolve_entry + + entries = [ + FedEntry("t", "guild", "", "guild-def"), + FedEntry("t", "channel", "c1", "channel-def"), + ] + result = _resolve_entry(entries, "c1", "u1") + assert result is not None + assert result.definition == "channel-def" + + +def test_resolve_entry_returns_none_when_no_match() -> None: + from lang2sql.tools.semantic_federation import _resolve_entry + + entries = [FedEntry("t", "channel", "other-channel", "def")] + assert _resolve_entry(entries, "c1", "u1") is None diff --git a/tests/test_setup_wizard.py b/tests/test_setup_wizard.py index e6019ca..f9a7f08 100644 --- a/tests/test_setup_wizard.py +++ b/tests/test_setup_wizard.py @@ -18,37 +18,61 @@ from lang2sql.frontends.discord.commands import CommandHandlers from lang2sql.tenancy.concierge import ContextConcierge - # --- dsn_builder --------------------------------------------------------- + def test_assemble_postgres_url(): - spec = assemble("postgresql", { - "host": "db.example.com", "port": "5432", "database": "analytics", - "user": "u", "password": "p", - }) + spec = assemble( + "postgresql", + { + "host": "db.example.com", + "port": "5432", + "database": "analytics", + "user": "u", + "password": "p", + }, + ) assert spec.dsn == "postgresql+psycopg://u:p@db.example.com:5432/analytics" assert spec.extras == {} def test_assemble_url_encodes_special_chars_in_password(): - spec = assemble("postgresql", { - "host": "h", "port": "5432", "database": "d", "user": "u", "password": "p@ss/w:rd", - }) + spec = assemble( + "postgresql", + { + "host": "h", + "port": "5432", + "database": "d", + "user": "u", + "password": "p@ss/w:rd", + }, + ) assert "p%40ss%2Fw%3Ard" in spec.dsn # @, /, : all encoded def test_assemble_snowflake_attaches_warehouse(): - spec = assemble("snowflake", { - "account": "ab12345.us-east-1", "user": "u", "password": "p", - "database": "DB", "warehouse": "WH", - }) + spec = assemble( + "snowflake", + { + "account": "ab12345.us-east-1", + "user": "u", + "password": "p", + "database": "DB", + "warehouse": "WH", + }, + ) assert "warehouse=WH" in spec.dsn and "@ab12345.us-east-1/DB" in spec.dsn def test_assemble_d1_returns_token_in_extras(): - spec = assemble("d1", { - "account_id": "acct", "database_id": "db", "api_token": "secret", - }) + spec = assemble( + "d1", + { + "account_id": "acct", + "database_id": "db", + "api_token": "secret", + }, + ) assert spec.dsn == "d1://acct/db" assert spec.extras == {"d1_token": "secret"} @@ -65,8 +89,10 @@ def test_assemble_unknown_db_type_raises(): # --- register_db_for_guild end-to-end (real sqlite) ---------------------- + def _seed_sqlite(path: str) -> None: from sqlalchemy import create_engine, text + eng = create_engine(f"sqlite:///{path}") with eng.begin() as conn: conn.execute(text("CREATE TABLE products (id INTEGER PRIMARY KEY, name TEXT)")) @@ -109,10 +135,19 @@ def test_register_db_for_guild_unknown_driver_gives_friendly_error(): identity = Identity(user_id="u", guild_id="g-x", channel_id="c") # Snowflake driver isn't installed in this env; the handler should catch # ModuleNotFoundError and produce a clear, non-technical message. - res = asyncio.run(handlers.register_db_for_guild( - identity, "snowflake", - {"account": "a", "user": "u", "password": "p", "database": "d", "warehouse": "w"}, - )) + res = asyncio.run( + handlers.register_db_for_guild( + identity, + "snowflake", + { + "account": "a", + "user": "u", + "password": "p", + "database": "d", + "warehouse": "w", + }, + ) + ) assert "uv sync --extra snowflake" in res.text or "Couldn't connect" in res.text @@ -120,14 +155,19 @@ def test_register_db_for_guild_missing_field_reports_setup_error(): concierge = ContextConcierge() handlers = CommandHandlers(concierge) identity = Identity(user_id="u", guild_id="g", channel_id="c") - res = asyncio.run(handlers.register_db_for_guild( - identity, "postgresql", {"host": "h"}, # missing user/password/db - )) + res = asyncio.run( + handlers.register_db_for_guild( + identity, + "postgresql", + {"host": "h"}, # missing user/password/db + ) + ) assert "Setup error" in res.text and "missing required" in res.text # --- concierge per-scope explorer routing -------------------------------- + def test_concierge_per_scope_dsn_routes_correctly(tmp_path): db = tmp_path / "scoped.db" _seed_sqlite(str(db)) @@ -180,6 +220,7 @@ def test_forget_explorer_busts_the_cache(tmp_path): # --- UI module import smoke ---------------------------------------------- + def test_setup_wizard_module_imports_without_discord_runtime(): # The wizard imports discord.ui at module level. Make sure that succeeds in # a no-gateway environment — the same contract as bot.py's import-safety. diff --git a/tests/test_tenancy.py b/tests/test_tenancy.py index a6afaa2..4bb097d 100644 --- a/tests/test_tenancy.py +++ b/tests/test_tenancy.py @@ -41,7 +41,9 @@ def test_build_context_populates_llm_and_session() -> None: try: concierge = ContextConcierge() identity = Identity(user_id="u1", guild_id="g", channel_id="c") - ctx = asyncio.run(concierge.build_context(identity, user_text="how many orders?")) + ctx = asyncio.run( + concierge.build_context(identity, user_text="how many orders?") + ) assert isinstance(ctx, HarnessContext) assert isinstance(ctx.llm, FakeLLM) # no key → fallback