diff --git a/bench/ecommerce_demo.py b/bench/ecommerce_demo.py index 0fab70d..0ebb29f 100644 --- a/bench/ecommerce_demo.py +++ b/bench/ecommerce_demo.py @@ -28,7 +28,13 @@ 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, + _render_effective, + _load_all, + _resolve_term, +) # Stable IDs for the demo guild and its two channels. GUILD = "acme-shop" @@ -50,7 +56,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 +101,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 +114,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 +149,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/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/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/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..628bd32 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) diff --git a/src/lang2sql/tools/__init__.py b/src/lang2sql/tools/__init__.py index b8d0813..726b2c4 100644 --- a/src/lang2sql/tools/__init__.py +++ b/src/lang2sql/tools/__init__.py @@ -23,8 +23,14 @@ __all__ = [ "build_default_tools", - "RunSQL", "ExploreSchema", "EnrichSchema", "SemanticFederationTool", - "OrgSetupTool", "Remember", "AskUser", "IngestDoc", + "RunSQL", + "ExploreSchema", + "EnrichSchema", + "SemanticFederationTool", + "OrgSetupTool", + "Remember", + "AskUser", + "IngestDoc", ] 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..adc5ebf 100644 --- a/src/lang2sql/tools/ingest_doc.py +++ b/src/lang2sql/tools/ingest_doc.py @@ -40,8 +40,14 @@ def spec(self) -> ToolSpec: 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)", + }, }, }, ) @@ -51,11 +57,19 @@ async def run(self, args: dict[str, Any], ctx: "HarnessContext") -> ToolResult: 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) + return ToolResult( + call_id="", + content="provide a document 'ref' or inline 'content'", + is_error=True, + ) - 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." + ) lines = ["Proposed definitions (confirm to register):"] for c in candidates: 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..f58a297 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 + ":") @@ -337,10 +414,7 @@ def build_prompt_section(store: Any, scope: str, channel_id: str, user_id: str) if line: lines.append(line) - header = ( - "## Business Terminology\n" - "(lookup 우선순위: 개인 > 채널(팀) > 전사)\n" - ) + header = "## Business Terminology\n" "(lookup 우선순위: 개인 > 채널(팀) > 전사)\n" body = "\n".join(lines) if lines else "(없음)" return header + body + "\n\n" + _AMBIGUOUS_TERM_POLICY @@ -360,7 +434,10 @@ 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: 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_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..d1295a6 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -26,7 +26,16 @@ 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", + } def test_run_sql_passes_gate_and_returns_rows(): @@ -50,11 +59,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 +79,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_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..7f8b452 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,53 @@ 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 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