234 lines
8.9 KiB
Python
234 lines
8.9 KiB
Python
"""Ask — natural-language SQL over the SQL-backed vector targets.
|
|
|
|
Vanna 2.0 (vanna-ai/vanna, agent-based rewrite) drives an Ollama LLM
|
|
with a RunSqlTool pointed at the same Postgres/MariaDB servers Thicket
|
|
ingests into. The question is prefixed with the corpus DDL so the
|
|
model writes correct SQL; only SELECT-style reads are requested.
|
|
|
|
Connections follow the same Unix conventions as the vector stores:
|
|
PGHOST/PGPORT/PGUSER/PGPASSWORD/PGDATABASE (or PGDSN) for pgvector,
|
|
MARIADB_HOST/PORT/USER/PASSWORD/DATABASE for mariadb.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import os
|
|
import tempfile
|
|
from collections.abc import Callable
|
|
|
|
# Targets whose corpus lives in a SQL database Vanna can query.
|
|
SQL_TARGETS = ("pgvector", "mariadb")
|
|
|
|
|
|
class AskUnavailable(Exception):
|
|
"""Raised when the ask interaction cannot run (target, deps)."""
|
|
|
|
|
|
def _pg_params() -> dict:
|
|
if os.environ.get("PGDSN"):
|
|
return {"connection_string": os.environ["PGDSN"]}
|
|
return {
|
|
"host": os.environ.get("PGHOST", "localhost"),
|
|
"port": int(os.environ.get("PGPORT", "5432")),
|
|
"database": os.environ.get("PGDATABASE", "thicket"),
|
|
"user": os.environ.get("PGUSER", "thicket"),
|
|
"password": os.environ.get("PGPASSWORD", ""),
|
|
}
|
|
|
|
|
|
def _mariadb_params() -> dict:
|
|
return {
|
|
"host": os.environ.get("MARIADB_HOST", "127.0.0.1"),
|
|
"port": int(os.environ.get("MARIADB_PORT", "3306")),
|
|
"database": os.environ.get("MARIADB_DATABASE", "thicket"),
|
|
"user": os.environ.get("MARIADB_USER", "root"),
|
|
"password": os.environ.get("MARIADB_PASSWORD", ""),
|
|
}
|
|
|
|
|
|
_RUNNER_PARAMS = {"pgvector": _pg_params, "mariadb": _mariadb_params}
|
|
|
|
|
|
def target_ddl(target: str, collection: str = "second_brain") -> str:
|
|
"""Corpus DDL + column documentation for the question context —
|
|
the exact shape Thicket's ensure_collection creates."""
|
|
match target:
|
|
case "pgvector":
|
|
return (
|
|
f'CREATE TABLE "{collection}" ('
|
|
"id TEXT PRIMARY KEY, doc_key TEXT NOT NULL, "
|
|
"embedding vector(384), payload JSONB);"
|
|
"\n-- payload fields: document_title text, obsidian_path text,"
|
|
" section_header text, content text, chunk_index int, doc_key text"
|
|
)
|
|
case "mariadb":
|
|
table = "".join(c if c.isalnum() or c == "_" else "_" for c in collection)
|
|
return (
|
|
f"CREATE TABLE `{table}` ("
|
|
"id VARCHAR(36) PRIMARY KEY, doc_key VARCHAR(64) NOT NULL, "
|
|
"embedding VECTOR(384) NOT NULL, payload JSON);"
|
|
"\n-- payload fields: document_title, obsidian_path,"
|
|
" section_header, content, chunk_index, doc_key"
|
|
)
|
|
raise AskUnavailable(
|
|
f"target '{target}' has no SQL corpus — ask works with: "
|
|
f"{', '.join(SQL_TARGETS)}"
|
|
)
|
|
|
|
|
|
def _component_text(component) -> str | None:
|
|
"""Text from a yielded component — vanna wraps each Rich component
|
|
(RichText, DataFrame, status pings) in a UiComponent envelope."""
|
|
rich = getattr(component, "rich_component", None) or component
|
|
for attr in ("content", "text", "markdown", "value"):
|
|
value = getattr(rich, attr, None)
|
|
if isinstance(value, str) and value.strip():
|
|
return value.strip()
|
|
df = getattr(rich, "df", None)
|
|
if df is not None and hasattr(df, "to_string"):
|
|
return df.head(20).to_string()
|
|
return None
|
|
|
|
|
|
class VannaAsker:
|
|
"""One question at a time against the corpus database."""
|
|
|
|
def __init__(self, target: str, collection: str, llm_model: str,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
self._target = target
|
|
if target not in SQL_TARGETS:
|
|
raise AskUnavailable(
|
|
f"ask needs a SQL-backed target ({', '.join(SQL_TARGETS)}) — "
|
|
f"current target: '{target}'"
|
|
)
|
|
try:
|
|
from vanna import Agent, AgentConfig
|
|
from vanna.core.registry import ToolRegistry
|
|
from vanna.core.user import RequestContext, User, UserResolver
|
|
from vanna.integrations.ollama import OllamaLlmService
|
|
from vanna.tools import RunSqlTool
|
|
except ImportError as e:
|
|
raise AskUnavailable(
|
|
"vanna not installed — run: pip install 'thicket[ask]'"
|
|
) from e
|
|
|
|
class _LocalUserResolver(UserResolver):
|
|
"""Single-user resolver: every ask is the same local user."""
|
|
async def resolve_user(self, request_context) -> User:
|
|
return User(id="thicket-local", username="thicket",
|
|
email="thicket@local", group_memberships=["user"])
|
|
|
|
from vanna.capabilities.agent_memory.base import AgentMemory
|
|
|
|
class _StatelessMemory(AgentMemory):
|
|
"""No persistence between asks — every question stands alone."""
|
|
|
|
def save_text_memory(self, content, context):
|
|
return None
|
|
|
|
def save_tool_usage(self, question, tool_name, args, context,
|
|
success=True, metadata=None):
|
|
return None
|
|
|
|
def get_recent_memories(self, context, limit=10):
|
|
return []
|
|
|
|
def get_recent_text_memories(self, context, limit=10):
|
|
return []
|
|
|
|
def search_similar_usage(self, question, context, *, limit=10,
|
|
similarity_threshold=0.7,
|
|
tool_name_filter=None):
|
|
return []
|
|
|
|
def search_text_memories(self, query, context, *, limit=10,
|
|
similarity_threshold=0.7):
|
|
return []
|
|
|
|
def clear_memories(self, context, tool_name=None, before_date=None):
|
|
return 0
|
|
|
|
def delete_by_id(self, context, memory_id):
|
|
return False
|
|
|
|
def delete_text_memory(self, context, memory_id):
|
|
return False
|
|
|
|
runner_cls = self._runner_cls(target)
|
|
runner = runner_cls(**_RUNNER_PARAMS[target]())
|
|
|
|
# Scope the tool's result-CSV scratch to a temp dir — without
|
|
# this, every ask litters a hash-named folder in the cwd.
|
|
from vanna.tools import LocalFileSystem
|
|
|
|
registry = ToolRegistry()
|
|
registry.register_local_tool(
|
|
RunSqlTool(
|
|
sql_runner=runner,
|
|
file_system=LocalFileSystem(
|
|
working_directory=tempfile.mkdtemp(prefix="thicket-ask-")),
|
|
),
|
|
access_groups=["user"])
|
|
|
|
resolver = _LocalUserResolver()
|
|
self._context = RequestContext(
|
|
remote_addr="127.0.0.1",
|
|
metadata={"source": "thicket"},
|
|
)
|
|
self._log = log
|
|
self._agent = Agent(
|
|
llm_service=OllamaLlmService(
|
|
model=llm_model, host=os.environ.get("OLLAMA_HOST"),
|
|
num_ctx=8192, temperature=0.1,
|
|
),
|
|
config=AgentConfig(stream_responses=False),
|
|
tool_registry=registry,
|
|
user_resolver=resolver,
|
|
agent_memory=_StatelessMemory(),
|
|
)
|
|
self._ddl = target_ddl(target, collection)
|
|
|
|
@staticmethod
|
|
def _runner_cls(target: str):
|
|
if target == "pgvector":
|
|
from vanna.integrations.postgres import PostgresRunner
|
|
return PostgresRunner
|
|
from vanna.integrations.mysql import MySQLRunner
|
|
return MySQLRunner
|
|
|
|
async def _ask_async(self, question: str) -> str:
|
|
json_hint = (
|
|
"payload->>'document_title'" if self._target == "pgvector"
|
|
else "JSON_UNQUOTE(JSON_EXTRACT(payload, '$.document_title'))"
|
|
)
|
|
prompt = (
|
|
"You are querying a knowledge-base corpus. Schema:\n"
|
|
f"{self._ddl}\n"
|
|
f"JSON fields are read with {json_hint}. "
|
|
"Read-only: SELECT queries only. "
|
|
"Answer with the final result only.\n\n"
|
|
f"Question: {question}"
|
|
)
|
|
# The stream carries status pings, reasoning, tool calls, and
|
|
# the final answer — the last RichText IS the answer.
|
|
texts: list[str] = []
|
|
async for component in self._agent.send_message(self._context, prompt):
|
|
text = _component_text(component)
|
|
if text:
|
|
texts.append(text)
|
|
return texts[-1] if texts else "(no answer produced)"
|
|
|
|
def ask(self, question: str) -> str:
|
|
"""Ask one question; returns the agent's answer text."""
|
|
self._log(f"Asking {self._agent.__class__.__name__} "
|
|
f"(target corpus, Ollama LLM)...")
|
|
answer = asyncio.run(self._ask_async(question))
|
|
return answer or "(no answer produced)"
|
|
|
|
|
|
def ask(question: str, target: str, collection: str = "second_brain",
|
|
llm_model: str = "llama3", log: Callable[[str], None] = lambda _msg: None) -> str:
|
|
"""Convenience one-shot entry for CLI and workers."""
|
|
return VannaAsker(target, collection, llm_model, log).ask(question)
|