Thicket/thicket/ask_vanna.py

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)