diff --git a/cmd_chat/agent/__init__.py b/cmd_chat/agent/__init__.py index 2d0578c..abd1f51 100644 --- a/cmd_chat/agent/__init__.py +++ b/cmd_chat/agent/__init__.py @@ -1,6 +1,6 @@ """hack-house AI agent bridge — model-agnostic agents that join a room.""" from .bridge import AgentBridge -from .providers import Msg, Provider, make_provider +from cmd_chat.ai.providers import Msg, Provider, make_provider __all__ = ["AgentBridge", "Msg", "Provider", "make_provider"] diff --git a/cmd_chat/agent/__main__.py b/cmd_chat/agent/__main__.py index 386e509..40ba885 100644 --- a/cmd_chat/agent/__main__.py +++ b/cmd_chat/agent/__main__.py @@ -31,8 +31,8 @@ import argparse import sys from .bridge import AgentBridge -from .profiles import load_profiles, provider_from_profile -from .providers import OllamaEmbedder, make_provider, preflight +from cmd_chat.ai.profiles import load_profiles, provider_from_profile +from cmd_chat.ai.providers import OllamaEmbedder, make_provider, preflight def _build_provider(args, ap): diff --git a/cmd_chat/agent/bridge.py b/cmd_chat/agent/bridge.py index 5fd50e3..cef2d47 100644 --- a/cmd_chat/agent/bridge.py +++ b/cmd_chat/agent/bridge.py @@ -22,7 +22,7 @@ import websockets from ..client.client import Client from .memory import MemoryIndex -from .providers import Msg, Provider, ToolsUnsupported +from cmd_chat.ai.providers import Msg, Provider, ToolsUnsupported DEFAULT_SYSTEM = ( "You are {name}, a helpful AI participant in an encrypted terminal chat " diff --git a/cmd_chat/agent/memory.py b/cmd_chat/agent/memory.py index c6a8cbb..a371e87 100644 --- a/cmd_chat/agent/memory.py +++ b/cmd_chat/agent/memory.py @@ -11,7 +11,7 @@ from __future__ import annotations import math from dataclasses import dataclass -from .providers import Msg +from cmd_chat.ai.providers import Msg @dataclass diff --git a/cmd_chat/agent/profiles.py b/cmd_chat/agent/profiles.py index 00dd466..7567f0c 100644 --- a/cmd_chat/agent/profiles.py +++ b/cmd_chat/agent/profiles.py @@ -1,102 +1,14 @@ -"""Named model profiles for the hack-house AI agent. +"""Backward-compatibility shim. -A *profile* maps a friendly name (``groq-llama``, ``local``, ``claude``) to a -provider + model + endpoint, so operators type ``--profile groq-llama`` instead -of remembering ``--provider openai --base-url … --model …``. This mirrors the -``models:`` list in Continue.dev and the ``model_list`` in a LiteLLM proxy: -each entry is ``{provider, model, base_url, api_key_env}``. - -Secrets are **never** stored here — ``api_key_env`` names an environment -variable to read the key from, keeping the file safe to commit and share. - -Lookup order (first hit wins): - 1. ``$HH_MODELS_FILE`` - 2. ``./models.toml`` (cwd) - 3. ``~/.config/hh/models.toml`` +Named model profiles moved to :mod:`cmd_chat.ai.profiles` so the operator bridge +and the ``/ai`` chat agent share one core. This module re-exports the public API +so existing imports (``from cmd_chat.agent.profiles import …``) keep working. +Prefer importing from ``cmd_chat.ai.profiles`` in new code. """ -from __future__ import annotations - -import os -from pathlib import Path - -try: # stdlib on 3.11+, falls back to the `tomli` backport on 3.10 - import tomllib -except ModuleNotFoundError: # pragma: no cover - import tomli as tomllib # type: ignore[no-redef] - -from .providers import Provider, make_provider - -_RECOGNIZED = {"provider", "model", "base_url", "host", "api_key_env", - "system", "context_window"} - - -def _candidate_paths(explicit: str | None) -> list[Path]: - if explicit: - return [Path(explicit).expanduser()] - paths = [] - env = os.environ.get("HH_MODELS_FILE") - if env: - paths.append(Path(env).expanduser()) - paths.append(Path.cwd() / "models.toml") - paths.append(Path.home() / ".config" / "hh" / "models.toml") - return paths - - -def find_profiles_file(explicit: str | None = None) -> Path | None: - for p in _candidate_paths(explicit): - if p.is_file(): - return p - return None - - -def load_profiles(explicit: str | None = None) -> dict[str, dict]: - """Return ``{name: profile_dict}`` from the first models.toml found.""" - path = find_profiles_file(explicit) - if path is None: - return {} - with path.open("rb") as fh: - data = tomllib.load(fh) - profiles: dict[str, dict] = {} - for name, body in data.items(): - if not isinstance(body, dict) or "provider" not in body: - continue # skip non-profile tables / malformed entries - unknown = set(body) - _RECOGNIZED - if unknown: - raise ValueError( - f"profile '{name}': unknown key(s) {', '.join(sorted(unknown))}" - ) - profiles[name] = body - return profiles - - -def provider_from_profile(prof: dict, *, name: str = "?", - model: str | None = None, - base_url: str | None = None) -> Provider: - """Build a :class:`Provider` from a profile dict. - - ``model`` / ``base_url`` (CLI flags) override the profile when given. The - api key is read from ``$`` and passed only to providers that - accept one, so an Ollama profile never sees a stray ``api_key`` kwarg. - """ - spec = prof["provider"] - custom = ":" in spec - opts: dict = {} - - mdl = model or prof.get("model") - bu = base_url or prof.get("base_url") - if bu and (spec == "openai" or custom): - opts["base_url"] = bu - if spec == "ollama" and prof.get("host"): - opts["host"] = prof["host"] - - key_env = prof.get("api_key_env") - if key_env and (spec in ("openai", "anthropic") or custom): - key = os.environ.get(key_env) - if not key: - raise SystemExit( - f"profile '{name}': ${key_env} is not set — export it first" - ) - opts["api_key"] = key - - return make_provider(spec, model=mdl, **opts) +from cmd_chat.ai.profiles import * # noqa: F401,F403 +from cmd_chat.ai.profiles import ( # noqa: F401 (explicit for star-safety) + find_profiles_file, + load_profiles, + provider_from_profile, +) diff --git a/cmd_chat/agent/providers.py b/cmd_chat/agent/providers.py index eb83fa2..0c27d3b 100644 --- a/cmd_chat/agent/providers.py +++ b/cmd_chat/agent/providers.py @@ -1,500 +1,21 @@ -"""Model-agnostic provider interface for the hack-house AI agent bridge. +"""Backward-compatibility shim. -A Provider turns a system prompt + conversation into a single reply string. -The bundled adapters speak plain HTTP via ``requests`` (already a dependency), -so no extra SDKs are required and any backend can be plugged in — including a -custom one via the ``module:Class`` spec. +The provider abstraction moved to :mod:`cmd_chat.ai.providers` so the operator +bridge and the ``/ai`` chat agent share one model-agnostic core. This module +re-exports the public API so existing imports +(``from cmd_chat.agent.providers import …``) keep working. Prefer importing from +``cmd_chat.ai.providers`` in new code. """ -from __future__ import annotations - -import importlib -import json -import os -import re -from dataclasses import dataclass -from typing import Protocol, runtime_checkable - -import requests - - -@dataclass -class Msg: - role: str # "system" | "user" | "assistant" - content: str - - -class ToolsUnsupported(RuntimeError): - """Raised by ``complete_with_tools`` when the backend model can't do function - calling — the native harness catches it and degrades to the simple injector.""" - - -@runtime_checkable -class Provider(Protocol): - name: str - model: str - - def complete(self, system: str, messages: list[Msg]) -> str: - ... - - # Optional: list models the backend can serve, for discovery/preflight. - # Providers that can't enumerate (e.g. a bespoke endpoint) may omit this. - def available_models(self) -> list[str]: - ... - - -class OllamaProvider: - """Local Ollama (default, recommended). No API key — privacy-preserving.""" - - name = "ollama" - - def __init__(self, model: str = "llama3", host: str | None = None, timeout: int = 240, - num_ctx: int = 4096, num_predict: int = 512, num_thread: int | None = None, - keep_alive: str = "30m"): - self.model = model - self.host = (host or os.environ.get("OLLAMA_HOST", "http://localhost:11434")).rstrip("/") - # Default 240s: the native tool-calling turn is NON-streaming, so on a - # contended CPU box a long write_file turn can exceed a tighter cap and - # surface as `[ai error: read timed out]`. Generous here, bounded loop above. - self.timeout = timeout - # On CPU, time-to-first-token is O(num_ctx) prefill, so keep the window - # modest (4096) rather than a GPU-mindset 8192. keep_alive pins the model - # so the next /ai doesn't pay a cold reload. num_thread defaults to - # Ollama's own (≈physical cores); set it explicitly to benchmark 4/6/8. - self.num_ctx = num_ctx - self.num_predict = num_predict - self.num_thread = num_thread - self.keep_alive = keep_alive - # Tri-state tool-calling capability cache: None=unprobed, True/False once a - # real /api/chat with `tools` either succeeds or is rejected by the model. - # The native harness reads this to skip retrying tools on a model that - # can't do them (and fall straight to the simple injector). - self._tools_ok: bool | None = None - - def _options(self, extra: dict | None = None) -> dict: - opts = {"num_ctx": self.num_ctx, "num_predict": self.num_predict} - if self.num_thread is not None: - opts["num_thread"] = self.num_thread - if extra: - opts.update(extra) - return opts - - def _raise_for_status(self, r: requests.Response) -> None: - """Turn an Ollama HTTP error into an actionable message. - - Ollama answers /api/chat with 404 + ``{"error": "model ... not found"}`` - when the model isn't pulled on the box running this agent. Because the - agent talks to *its own* localhost:11434, a teammate who summoned /ai - without that model pulled hits this even when the host has it. The bare - ``raise_for_status`` only reports "404 Not Found for url", hiding the - cause — so name the model, the host, and the fix instead. This text is - what the bridge posts to the room as ``[ai error: …]``. - """ - if r.ok: - return - try: - detail = (r.json().get("error") or "").strip() - except ValueError: - detail = (r.text or "").strip() - if r.status_code == 404: - raise RuntimeError( - f"model '{self.model}' isn't pulled on the ollama at {self.host} " - f"(the agent uses ollama on the machine that ran /ai, not the host). " - f"fix: `ollama pull {self.model}` there, or `/ai start ` " - f"for a cloud model. [{detail or 'model not found'}]" - ) - raise RuntimeError( - f"ollama at {self.host} returned {r.status_code}" - + (f": {detail}" if detail else "") - ) - - def complete(self, system: str, messages: list[Msg]) -> str: - payload = { - "model": self.model, - "stream": False, - "keep_alive": self.keep_alive, - "options": self._options(), - "messages": [{"role": "system", "content": system}] - + [{"role": m.role, "content": m.content} for m in messages], - } - r = requests.post(f"{self.host}/api/chat", json=payload, timeout=self.timeout) - self._raise_for_status(r) - return (r.json().get("message", {}).get("content") or "").strip() - - def supports_tools(self) -> bool | None: - """Cached tool-calling capability: None until the first ``complete_with_tools`` - call has either succeeded or been rejected by the model.""" - return self._tools_ok - - def complete_with_tools( - self, system: str, messages: list[dict], tools: list[dict] - ) -> tuple[str, list[dict], dict]: - """One non-streaming ``/api/chat`` turn carrying a ``tools`` schema. Used by - the native harness loop. ``messages`` are raw Ollama wire dicts (so the - caller can round-trip assistant ``tool_calls`` and ``tool`` results across - turns); ``system`` is prepended. Returns ``(text, tool_calls, usage)`` where - each call is ``{"name": str, "arguments": dict}`` and ``usage`` carries - Ollama's real token counts (``prompt_eval_count`` / ``eval_count``) so the - caller can budget context against TRUE tokens instead of a char estimate - (``{}`` if the server omits them). Raises ``ToolsUnsupported`` if the model - can't do function calling so the bridge can fall back to simple.""" - # Greedy decode (temperature 0) for the tool loop: at Ollama's default 0.8 a - # weak model "creatively" narrates the next step in prose or fabricates file - # content instead of emitting a deterministic structured call. The nudge loop - # changes the prompt between turns, so temp 0 still escapes a failing state on - # retry — it just stops sampling away from the correct tool-call format. This - # override is scoped to complete_with_tools; chat (complete/stream) keeps the - # model's default sampling so replies stay natural. - payload = { - "model": self.model, - "stream": False, - "keep_alive": self.keep_alive, - "options": self._options({"temperature": 0.0}), - "tools": tools, - "messages": [{"role": "system", "content": system}] + messages, - } - r = requests.post(f"{self.host}/api/chat", json=payload, timeout=self.timeout) - if not r.ok: - try: - detail = (r.json().get("error") or "").strip() - except ValueError: - detail = (r.text or "").strip() - if "does not support tools" in detail.lower(): - self._tools_ok = False - raise ToolsUnsupported(detail or f"{self.model} does not support tools") - self._raise_for_status(r) - self._tools_ok = True - data = r.json() - msg = data.get("message", {}) or {} - # Real token counts straight from Ollama — exact, free (already in the - # response), and used to calibrate the native loop's char-based estimate. - usage = {k: data[k] for k in ("prompt_eval_count", "eval_count") - if isinstance(data.get(k), int)} - text = (msg.get("content") or "").strip() - calls: list[dict] = [] - for tc in msg.get("tool_calls") or []: - fn = tc.get("function") or {} - args = fn.get("arguments") - if isinstance(args, str): - try: - args = json.loads(args) - except ValueError: - args = {} - calls.append({"name": fn.get("name", ""), "arguments": args or {}}) - # Small/quantized models (notably qwen2.5 on CPU) intermittently emit a valid - # tool call as literal text in `content` instead of the structured `tool_calls` - # field — qwen's `{…}`, but also bare/fenced JSON and - # alternate wrappers (``, ``). This is the single biggest - # score sink in the native-harness benchmark, so recover any well-formed JSON - # call here. Gate on the known tool names from `tools` so a stray JSON blob in - # prose can never be coerced into an action the model didn't structurally ask - # for. Only adopt the recovery when it actually found a call (a plain `DONE:` - # or prose turn is left untouched). - if not calls and text: - valid = {(t.get("function") or {}).get("name") for t in (tools or [])} - valid.discard(None) - recovered_text, recovered = self._extract_text_tool_calls(text, valid) - if recovered: - text, calls = recovered_text, recovered - return text, calls, usage - - # Wrapper tags a weak model wraps a leaked call (or its prose) in; stripped - # from the chat-facing text once the JSON inside is recovered. - _WRAP_TAGS = re.compile( - r"", - re.I, - ) - - # The SPLIT-form leak (qwen2.5:0.5b at temp 0, ~half its turns): the tool NAME in - # a `` tag and the arguments in a SEPARATE bare JSON object with no `name` - # key — `write_file{"path":…,"content":…}`. Captures the name; the - # decoder reads the args object that follows from the trailing `{`. - _NAMED_TAG = re.compile( - r"<(tool_call|tool_calls|function_call|function|tools)>\s*" - r"([a-zA-Z_]\w*)\s*\s*(?=\{)", - re.I, - ) - - @classmethod - def _coerce_call(cls, obj, valid_names) -> dict | None: - """Turn a decoded JSON object into a `{"name","arguments"}` call IF it - structurally is one for a KNOWN tool — else None. Unwraps the OpenAI-style - `{"function": {...}}` / `{"tool_call": {...}}` nesting and accepts either - `arguments` or qwen's `parameters` key. The `valid_names` gate is what makes - scanning arbitrary text safe: a random JSON blob in prose has no known tool - name, so it can never be coerced into an action.""" - if not isinstance(obj, dict): - return None - inner = obj.get("function") or obj.get("tool_call") - if isinstance(inner, dict): - obj = inner - name = obj.get("name") - if not isinstance(name, str) or not name: - return None - if valid_names and name not in valid_names: - return None - args = obj.get("arguments") - if args is None: - args = obj.get("parameters") - if isinstance(args, str): - try: - args = json.loads(args) - except ValueError: - args = {} - if not isinstance(args, dict): - args = {} - return {"name": name, "arguments": args} - - @classmethod - def _extract_text_tool_calls( - cls, text: str, valid_names: set | None = None - ) -> tuple[str, list[dict]]: - """Recover tool calls a small/quantized model emitted as TEXT in `content` - instead of the structured `tool_calls` field. Handles qwen's - `{json}` blocks plus the looser CPU-model leaks: bare - JSON, ```json fenced blocks, alternate wrapper tags (``, - ``), and the SPLIT form where the name sits in a tag and the - args follow as a separate object (`write_file{"path":…}`). - Scans for every JSON object via a decoder (so nested braces in arguments parse - correctly) and keeps ONLY those that resolve to a KNOWN tool — never freeform - prose, so it can't fabricate an action the model didn't structurally request. - Returns the text with the recovered JSON (and now-orphaned wrapper tags / code - fences) stripped, plus the calls.""" - dec = json.JSONDecoder() - calls: list[dict] = [] - spans: list[tuple[int, int]] = [] - # Split-form index: the `{` that opens an args object → (tool_name, tag_start), - # so the scan pairs that JSON as arguments and strips the whole tag+object. - split = {m.end(): (m.group(2), m.start()) for m in cls._NAMED_TAG.finditer(text)} - i, n = 0, len(text) - while i < n: - brace = text.find("{", i) - if brace == -1: - break - try: - obj, end = dec.raw_decode(text, brace) - except ValueError: - i = brace + 1 - continue - if brace in split: - # `NAME{args}` — name from the tag, this object is the args. - name, tag_start = split[brace] - if (not valid_names or name in valid_names) and isinstance(obj, dict): - calls.append({"name": name, "arguments": obj}) - spans.append((tag_start, end)) - else: - call = cls._coerce_call(obj, valid_names) - if call is not None: - calls.append(call) - spans.append((brace, end)) - i = end - if spans: - kept, last = [], 0 - for start, stop in spans: - kept.append(text[last:start]) - last = stop - kept.append(text[last:]) - text = "".join(kept) - # The JSON is gone; drop the wrapper tags and any now-empty code fences - # it sat in so the chat summary reads as clean prose. - text = cls._WRAP_TAGS.sub("", text) - text = re.sub(r"```[a-zA-Z]*\s*```", "", text) - text = re.sub(r"```[a-zA-Z]*|```", "", text) - text = text.strip() - return text, calls - - def stream(self, system: str, messages: list[Msg]): - """Yield reply text incrementally as Ollama generates it. On CPU the - perceived latency is TTFT, so streaming makes a slow reply feel live.""" - payload = { - "model": self.model, - "stream": True, - "keep_alive": self.keep_alive, - "options": self._options(), - "messages": [{"role": "system", "content": system}] - + [{"role": m.role, "content": m.content} for m in messages], - } - with requests.post(f"{self.host}/api/chat", json=payload, - timeout=self.timeout, stream=True) as r: - self._raise_for_status(r) - for line in r.iter_lines(): - if not line: - continue - chunk = json.loads(line) - piece = chunk.get("message", {}).get("content") - if piece: - yield piece - if chunk.get("done"): - break - - def available_models(self) -> list[str]: - r = requests.get(f"{self.host}/api/tags", timeout=self.timeout) - r.raise_for_status() - return [m.get("name", "") for m in r.json().get("models", [])] - - -class OllamaEmbedder: - """Local text embeddings via Ollama (default ``nomic-embed-text``), used for - the agent's in-RAM semantic recall. Local + free, so it stays on by default - regardless of which provider answers chat. No key, nothing persisted.""" - - name = "ollama-embed" - - def __init__(self, model: str = "nomic-embed-text", host: str | None = None, - timeout: int = 60, truncate_dim: int | None = 256): - self.model = model - self.host = (host or os.environ.get("OLLAMA_HOST", "http://localhost:11434")).rstrip("/") - self.timeout = timeout - # nomic-embed-text is Matryoshka (MRL)-trained, so its 768-dim vector can - # be truncated to a shorter prefix with little quality loss — faster - # pure-Python cosine and less RAM. Query + stored use the same dim, so - # cosine stays correct. None keeps the full vector. - self.truncate_dim = truncate_dim - - def embed(self, text: str) -> list[float]: - r = requests.post( - f"{self.host}/api/embeddings", - json={"model": self.model, "prompt": text}, - timeout=self.timeout, - ) - r.raise_for_status() - vec = r.json().get("embedding") or [] - if self.truncate_dim is not None: - vec = vec[: self.truncate_dim] - return vec - - -class AnthropicProvider: - """Anthropic Messages API. Cloud — opt-in. Needs ANTHROPIC_API_KEY.""" - - name = "anthropic" - - def __init__(self, model: str = "claude-opus-4-6", api_key: str | None = None, - timeout: int = 120, max_tokens: int = 1024): - self.model = model - self.api_key = api_key or os.environ.get("ANTHROPIC_API_KEY") - self.timeout = timeout - self.max_tokens = max_tokens - if not self.api_key: - raise ValueError("ANTHROPIC_API_KEY not set") - - def complete(self, system: str, messages: list[Msg]) -> str: - payload = { - "model": self.model, - "max_tokens": self.max_tokens, - "system": system, - "messages": [ - {"role": m.role, "content": m.content} - for m in messages - if m.role in ("user", "assistant") - ], - } - r = requests.post( - "https://api.anthropic.com/v1/messages", - json=payload, - timeout=self.timeout, - headers={ - "x-api-key": self.api_key, - "anthropic-version": "2023-06-01", - "content-type": "application/json", - }, - ) - r.raise_for_status() - blocks = r.json().get("content", []) - return "".join(b.get("text", "") for b in blocks).strip() - - def available_models(self) -> list[str]: - r = requests.get( - "https://api.anthropic.com/v1/models", - timeout=self.timeout, - headers={"x-api-key": self.api_key, "anthropic-version": "2023-06-01"}, - ) - r.raise_for_status() - return [m.get("id", "") for m in r.json().get("data", [])] - - -class OpenAICompatibleProvider: - """OpenAI-style /chat/completions — OpenAI, Groq, Together, local vLLM, etc.""" - - name = "openai" - - def __init__(self, model: str = "gpt-4o-mini", api_key: str | None = None, - base_url: str | None = None, timeout: int = 120): - self.model = model - self.api_key = api_key or os.environ.get("OPENAI_API_KEY", "") - self.base_url = (base_url or os.environ.get("OPENAI_BASE_URL", "https://api.openai.com/v1")).rstrip("/") - self.timeout = timeout - - def complete(self, system: str, messages: list[Msg]) -> str: - payload = { - "model": self.model, - "messages": [{"role": "system", "content": system}] - + [{"role": m.role, "content": m.content} for m in messages], - } - headers = {"content-type": "application/json"} - if self.api_key: - headers["authorization"] = f"Bearer {self.api_key}" - r = requests.post( - f"{self.base_url}/chat/completions", json=payload, headers=headers, timeout=self.timeout - ) - r.raise_for_status() - return r.json()["choices"][0]["message"]["content"].strip() - - def available_models(self) -> list[str]: - headers = {} - if self.api_key: - headers["authorization"] = f"Bearer {self.api_key}" - r = requests.get(f"{self.base_url}/models", headers=headers, timeout=self.timeout) - r.raise_for_status() - return [m.get("id", "") for m in r.json().get("data", [])] - - -_BUILTINS = { - "ollama": OllamaProvider, - "anthropic": AnthropicProvider, - "openai": OpenAICompatibleProvider, -} - - -def make_provider(spec: str, model: str | None = None, **opts) -> Provider: - """Build a provider. - - ``spec`` is a builtin name (``ollama`` / ``anthropic`` / ``openai``) or a - ``module:Class`` path to a custom Provider implementation. - """ - if ":" in spec: - mod_name, _, cls_name = spec.partition(":") - cls = getattr(importlib.import_module(mod_name), cls_name) - else: - cls = _BUILTINS.get(spec) - if cls is None: - raise ValueError(f"unknown provider '{spec}' (builtins: {', '.join(_BUILTINS)})") - if model is not None: - opts["model"] = model - return cls(**opts) - - -def preflight(provider: Provider) -> tuple[bool, str]: - """Cheap reachability + model-presence check before joining a room. - - Returns ``(ok, message)``. Lets ``/ai start`` fail fast with a clear reason - (backend down / model not pulled / key missing) instead of erroring on the - first question. Providers without ``available_models`` are assumed reachable. - """ - discover = getattr(provider, "available_models", None) - if discover is None: - return True, f"{provider.name}: no discovery endpoint — assuming reachable" - try: - models = discover() - except Exception as e: # noqa: BLE001 — any failure means "not reachable yet" - return False, f"{provider.name}: cannot reach backend ({e})" - if provider.model in models: - return True, f"{provider.name}/{provider.model}: reachable" - if models: - sample = ", ".join(models[:8]) - more = "…" if len(models) > 8 else "" - return False, ( - f"{provider.name}: model '{provider.model}' not available. " - f"reachable models: {sample}{more}" - ) - return True, f"{provider.name}: reachable (empty model list — skipping check)" +from cmd_chat.ai.providers import * # noqa: F401,F403 +from cmd_chat.ai.providers import ( # noqa: F401 (explicit for star-safety) + Msg, + OllamaEmbedder, + OllamaProvider, + AnthropicProvider, + OpenAICompatibleProvider, + Provider, + ToolsUnsupported, + make_provider, + preflight, +) diff --git a/cmd_chat/ai/__init__.py b/cmd_chat/ai/__init__.py new file mode 100644 index 0000000..0e43438 --- /dev/null +++ b/cmd_chat/ai/__init__.py @@ -0,0 +1,37 @@ +"""Shared model-agnostic AI core for hack-house. + +Houses the provider abstraction (``providers``) and named model profiles +(``profiles``) consumed by BOTH the in-room ``/ai`` chat agent and the operator +bridge. Hoisted here from ``cmd_chat.agent`` so the operator can run any +function-calling model as an operator, not just Claude. The old +``cmd_chat.agent.{providers,profiles}`` paths remain as thin re-export shims for +backward compatibility. +""" + +from .providers import ( + Msg, + OllamaEmbedder, + OllamaProvider, + Provider, + ToolsUnsupported, + make_provider, + preflight, +) +from .profiles import ( + find_profiles_file, + load_profiles, + provider_from_profile, +) + +__all__ = [ + "Msg", + "Provider", + "ToolsUnsupported", + "OllamaProvider", + "OllamaEmbedder", + "make_provider", + "preflight", + "load_profiles", + "find_profiles_file", + "provider_from_profile", +] diff --git a/cmd_chat/ai/profiles.py b/cmd_chat/ai/profiles.py new file mode 100644 index 0000000..00dd466 --- /dev/null +++ b/cmd_chat/ai/profiles.py @@ -0,0 +1,102 @@ +"""Named model profiles for the hack-house AI agent. + +A *profile* maps a friendly name (``groq-llama``, ``local``, ``claude``) to a +provider + model + endpoint, so operators type ``--profile groq-llama`` instead +of remembering ``--provider openai --base-url … --model …``. This mirrors the +``models:`` list in Continue.dev and the ``model_list`` in a LiteLLM proxy: +each entry is ``{provider, model, base_url, api_key_env}``. + +Secrets are **never** stored here — ``api_key_env`` names an environment +variable to read the key from, keeping the file safe to commit and share. + +Lookup order (first hit wins): + 1. ``$HH_MODELS_FILE`` + 2. ``./models.toml`` (cwd) + 3. ``~/.config/hh/models.toml`` +""" + +from __future__ import annotations + +import os +from pathlib import Path + +try: # stdlib on 3.11+, falls back to the `tomli` backport on 3.10 + import tomllib +except ModuleNotFoundError: # pragma: no cover + import tomli as tomllib # type: ignore[no-redef] + +from .providers import Provider, make_provider + +_RECOGNIZED = {"provider", "model", "base_url", "host", "api_key_env", + "system", "context_window"} + + +def _candidate_paths(explicit: str | None) -> list[Path]: + if explicit: + return [Path(explicit).expanduser()] + paths = [] + env = os.environ.get("HH_MODELS_FILE") + if env: + paths.append(Path(env).expanduser()) + paths.append(Path.cwd() / "models.toml") + paths.append(Path.home() / ".config" / "hh" / "models.toml") + return paths + + +def find_profiles_file(explicit: str | None = None) -> Path | None: + for p in _candidate_paths(explicit): + if p.is_file(): + return p + return None + + +def load_profiles(explicit: str | None = None) -> dict[str, dict]: + """Return ``{name: profile_dict}`` from the first models.toml found.""" + path = find_profiles_file(explicit) + if path is None: + return {} + with path.open("rb") as fh: + data = tomllib.load(fh) + profiles: dict[str, dict] = {} + for name, body in data.items(): + if not isinstance(body, dict) or "provider" not in body: + continue # skip non-profile tables / malformed entries + unknown = set(body) - _RECOGNIZED + if unknown: + raise ValueError( + f"profile '{name}': unknown key(s) {', '.join(sorted(unknown))}" + ) + profiles[name] = body + return profiles + + +def provider_from_profile(prof: dict, *, name: str = "?", + model: str | None = None, + base_url: str | None = None) -> Provider: + """Build a :class:`Provider` from a profile dict. + + ``model`` / ``base_url`` (CLI flags) override the profile when given. The + api key is read from ``$`` and passed only to providers that + accept one, so an Ollama profile never sees a stray ``api_key`` kwarg. + """ + spec = prof["provider"] + custom = ":" in spec + opts: dict = {} + + mdl = model or prof.get("model") + bu = base_url or prof.get("base_url") + if bu and (spec == "openai" or custom): + opts["base_url"] = bu + if spec == "ollama" and prof.get("host"): + opts["host"] = prof["host"] + + key_env = prof.get("api_key_env") + if key_env and (spec in ("openai", "anthropic") or custom): + key = os.environ.get(key_env) + if not key: + raise SystemExit( + f"profile '{name}': ${key_env} is not set — export it first" + ) + opts["api_key"] = key + + return make_provider(spec, model=mdl, **opts) diff --git a/cmd_chat/ai/providers.py b/cmd_chat/ai/providers.py new file mode 100644 index 0000000..eb83fa2 --- /dev/null +++ b/cmd_chat/ai/providers.py @@ -0,0 +1,500 @@ +"""Model-agnostic provider interface for the hack-house AI agent bridge. + +A Provider turns a system prompt + conversation into a single reply string. +The bundled adapters speak plain HTTP via ``requests`` (already a dependency), +so no extra SDKs are required and any backend can be plugged in — including a +custom one via the ``module:Class`` spec. +""" + +from __future__ import annotations + +import importlib +import json +import os +import re +from dataclasses import dataclass +from typing import Protocol, runtime_checkable + +import requests + + +@dataclass +class Msg: + role: str # "system" | "user" | "assistant" + content: str + + +class ToolsUnsupported(RuntimeError): + """Raised by ``complete_with_tools`` when the backend model can't do function + calling — the native harness catches it and degrades to the simple injector.""" + + +@runtime_checkable +class Provider(Protocol): + name: str + model: str + + def complete(self, system: str, messages: list[Msg]) -> str: + ... + + # Optional: list models the backend can serve, for discovery/preflight. + # Providers that can't enumerate (e.g. a bespoke endpoint) may omit this. + def available_models(self) -> list[str]: + ... + + +class OllamaProvider: + """Local Ollama (default, recommended). No API key — privacy-preserving.""" + + name = "ollama" + + def __init__(self, model: str = "llama3", host: str | None = None, timeout: int = 240, + num_ctx: int = 4096, num_predict: int = 512, num_thread: int | None = None, + keep_alive: str = "30m"): + self.model = model + self.host = (host or os.environ.get("OLLAMA_HOST", "http://localhost:11434")).rstrip("/") + # Default 240s: the native tool-calling turn is NON-streaming, so on a + # contended CPU box a long write_file turn can exceed a tighter cap and + # surface as `[ai error: read timed out]`. Generous here, bounded loop above. + self.timeout = timeout + # On CPU, time-to-first-token is O(num_ctx) prefill, so keep the window + # modest (4096) rather than a GPU-mindset 8192. keep_alive pins the model + # so the next /ai doesn't pay a cold reload. num_thread defaults to + # Ollama's own (≈physical cores); set it explicitly to benchmark 4/6/8. + self.num_ctx = num_ctx + self.num_predict = num_predict + self.num_thread = num_thread + self.keep_alive = keep_alive + # Tri-state tool-calling capability cache: None=unprobed, True/False once a + # real /api/chat with `tools` either succeeds or is rejected by the model. + # The native harness reads this to skip retrying tools on a model that + # can't do them (and fall straight to the simple injector). + self._tools_ok: bool | None = None + + def _options(self, extra: dict | None = None) -> dict: + opts = {"num_ctx": self.num_ctx, "num_predict": self.num_predict} + if self.num_thread is not None: + opts["num_thread"] = self.num_thread + if extra: + opts.update(extra) + return opts + + def _raise_for_status(self, r: requests.Response) -> None: + """Turn an Ollama HTTP error into an actionable message. + + Ollama answers /api/chat with 404 + ``{"error": "model ... not found"}`` + when the model isn't pulled on the box running this agent. Because the + agent talks to *its own* localhost:11434, a teammate who summoned /ai + without that model pulled hits this even when the host has it. The bare + ``raise_for_status`` only reports "404 Not Found for url", hiding the + cause — so name the model, the host, and the fix instead. This text is + what the bridge posts to the room as ``[ai error: …]``. + """ + if r.ok: + return + try: + detail = (r.json().get("error") or "").strip() + except ValueError: + detail = (r.text or "").strip() + if r.status_code == 404: + raise RuntimeError( + f"model '{self.model}' isn't pulled on the ollama at {self.host} " + f"(the agent uses ollama on the machine that ran /ai, not the host). " + f"fix: `ollama pull {self.model}` there, or `/ai start ` " + f"for a cloud model. [{detail or 'model not found'}]" + ) + raise RuntimeError( + f"ollama at {self.host} returned {r.status_code}" + + (f": {detail}" if detail else "") + ) + + def complete(self, system: str, messages: list[Msg]) -> str: + payload = { + "model": self.model, + "stream": False, + "keep_alive": self.keep_alive, + "options": self._options(), + "messages": [{"role": "system", "content": system}] + + [{"role": m.role, "content": m.content} for m in messages], + } + r = requests.post(f"{self.host}/api/chat", json=payload, timeout=self.timeout) + self._raise_for_status(r) + return (r.json().get("message", {}).get("content") or "").strip() + + def supports_tools(self) -> bool | None: + """Cached tool-calling capability: None until the first ``complete_with_tools`` + call has either succeeded or been rejected by the model.""" + return self._tools_ok + + def complete_with_tools( + self, system: str, messages: list[dict], tools: list[dict] + ) -> tuple[str, list[dict], dict]: + """One non-streaming ``/api/chat`` turn carrying a ``tools`` schema. Used by + the native harness loop. ``messages`` are raw Ollama wire dicts (so the + caller can round-trip assistant ``tool_calls`` and ``tool`` results across + turns); ``system`` is prepended. Returns ``(text, tool_calls, usage)`` where + each call is ``{"name": str, "arguments": dict}`` and ``usage`` carries + Ollama's real token counts (``prompt_eval_count`` / ``eval_count``) so the + caller can budget context against TRUE tokens instead of a char estimate + (``{}`` if the server omits them). Raises ``ToolsUnsupported`` if the model + can't do function calling so the bridge can fall back to simple.""" + # Greedy decode (temperature 0) for the tool loop: at Ollama's default 0.8 a + # weak model "creatively" narrates the next step in prose or fabricates file + # content instead of emitting a deterministic structured call. The nudge loop + # changes the prompt between turns, so temp 0 still escapes a failing state on + # retry — it just stops sampling away from the correct tool-call format. This + # override is scoped to complete_with_tools; chat (complete/stream) keeps the + # model's default sampling so replies stay natural. + payload = { + "model": self.model, + "stream": False, + "keep_alive": self.keep_alive, + "options": self._options({"temperature": 0.0}), + "tools": tools, + "messages": [{"role": "system", "content": system}] + messages, + } + r = requests.post(f"{self.host}/api/chat", json=payload, timeout=self.timeout) + if not r.ok: + try: + detail = (r.json().get("error") or "").strip() + except ValueError: + detail = (r.text or "").strip() + if "does not support tools" in detail.lower(): + self._tools_ok = False + raise ToolsUnsupported(detail or f"{self.model} does not support tools") + self._raise_for_status(r) + self._tools_ok = True + data = r.json() + msg = data.get("message", {}) or {} + # Real token counts straight from Ollama — exact, free (already in the + # response), and used to calibrate the native loop's char-based estimate. + usage = {k: data[k] for k in ("prompt_eval_count", "eval_count") + if isinstance(data.get(k), int)} + text = (msg.get("content") or "").strip() + calls: list[dict] = [] + for tc in msg.get("tool_calls") or []: + fn = tc.get("function") or {} + args = fn.get("arguments") + if isinstance(args, str): + try: + args = json.loads(args) + except ValueError: + args = {} + calls.append({"name": fn.get("name", ""), "arguments": args or {}}) + # Small/quantized models (notably qwen2.5 on CPU) intermittently emit a valid + # tool call as literal text in `content` instead of the structured `tool_calls` + # field — qwen's `{…}`, but also bare/fenced JSON and + # alternate wrappers (``, ``). This is the single biggest + # score sink in the native-harness benchmark, so recover any well-formed JSON + # call here. Gate on the known tool names from `tools` so a stray JSON blob in + # prose can never be coerced into an action the model didn't structurally ask + # for. Only adopt the recovery when it actually found a call (a plain `DONE:` + # or prose turn is left untouched). + if not calls and text: + valid = {(t.get("function") or {}).get("name") for t in (tools or [])} + valid.discard(None) + recovered_text, recovered = self._extract_text_tool_calls(text, valid) + if recovered: + text, calls = recovered_text, recovered + return text, calls, usage + + # Wrapper tags a weak model wraps a leaked call (or its prose) in; stripped + # from the chat-facing text once the JSON inside is recovered. + _WRAP_TAGS = re.compile( + r"", + re.I, + ) + + # The SPLIT-form leak (qwen2.5:0.5b at temp 0, ~half its turns): the tool NAME in + # a `` tag and the arguments in a SEPARATE bare JSON object with no `name` + # key — `write_file{"path":…,"content":…}`. Captures the name; the + # decoder reads the args object that follows from the trailing `{`. + _NAMED_TAG = re.compile( + r"<(tool_call|tool_calls|function_call|function|tools)>\s*" + r"([a-zA-Z_]\w*)\s*\s*(?=\{)", + re.I, + ) + + @classmethod + def _coerce_call(cls, obj, valid_names) -> dict | None: + """Turn a decoded JSON object into a `{"name","arguments"}` call IF it + structurally is one for a KNOWN tool — else None. Unwraps the OpenAI-style + `{"function": {...}}` / `{"tool_call": {...}}` nesting and accepts either + `arguments` or qwen's `parameters` key. The `valid_names` gate is what makes + scanning arbitrary text safe: a random JSON blob in prose has no known tool + name, so it can never be coerced into an action.""" + if not isinstance(obj, dict): + return None + inner = obj.get("function") or obj.get("tool_call") + if isinstance(inner, dict): + obj = inner + name = obj.get("name") + if not isinstance(name, str) or not name: + return None + if valid_names and name not in valid_names: + return None + args = obj.get("arguments") + if args is None: + args = obj.get("parameters") + if isinstance(args, str): + try: + args = json.loads(args) + except ValueError: + args = {} + if not isinstance(args, dict): + args = {} + return {"name": name, "arguments": args} + + @classmethod + def _extract_text_tool_calls( + cls, text: str, valid_names: set | None = None + ) -> tuple[str, list[dict]]: + """Recover tool calls a small/quantized model emitted as TEXT in `content` + instead of the structured `tool_calls` field. Handles qwen's + `{json}` blocks plus the looser CPU-model leaks: bare + JSON, ```json fenced blocks, alternate wrapper tags (``, + ``), and the SPLIT form where the name sits in a tag and the + args follow as a separate object (`write_file{"path":…}`). + Scans for every JSON object via a decoder (so nested braces in arguments parse + correctly) and keeps ONLY those that resolve to a KNOWN tool — never freeform + prose, so it can't fabricate an action the model didn't structurally request. + Returns the text with the recovered JSON (and now-orphaned wrapper tags / code + fences) stripped, plus the calls.""" + dec = json.JSONDecoder() + calls: list[dict] = [] + spans: list[tuple[int, int]] = [] + # Split-form index: the `{` that opens an args object → (tool_name, tag_start), + # so the scan pairs that JSON as arguments and strips the whole tag+object. + split = {m.end(): (m.group(2), m.start()) for m in cls._NAMED_TAG.finditer(text)} + i, n = 0, len(text) + while i < n: + brace = text.find("{", i) + if brace == -1: + break + try: + obj, end = dec.raw_decode(text, brace) + except ValueError: + i = brace + 1 + continue + if brace in split: + # `NAME{args}` — name from the tag, this object is the args. + name, tag_start = split[brace] + if (not valid_names or name in valid_names) and isinstance(obj, dict): + calls.append({"name": name, "arguments": obj}) + spans.append((tag_start, end)) + else: + call = cls._coerce_call(obj, valid_names) + if call is not None: + calls.append(call) + spans.append((brace, end)) + i = end + if spans: + kept, last = [], 0 + for start, stop in spans: + kept.append(text[last:start]) + last = stop + kept.append(text[last:]) + text = "".join(kept) + # The JSON is gone; drop the wrapper tags and any now-empty code fences + # it sat in so the chat summary reads as clean prose. + text = cls._WRAP_TAGS.sub("", text) + text = re.sub(r"```[a-zA-Z]*\s*```", "", text) + text = re.sub(r"```[a-zA-Z]*|```", "", text) + text = text.strip() + return text, calls + + def stream(self, system: str, messages: list[Msg]): + """Yield reply text incrementally as Ollama generates it. On CPU the + perceived latency is TTFT, so streaming makes a slow reply feel live.""" + payload = { + "model": self.model, + "stream": True, + "keep_alive": self.keep_alive, + "options": self._options(), + "messages": [{"role": "system", "content": system}] + + [{"role": m.role, "content": m.content} for m in messages], + } + with requests.post(f"{self.host}/api/chat", json=payload, + timeout=self.timeout, stream=True) as r: + self._raise_for_status(r) + for line in r.iter_lines(): + if not line: + continue + chunk = json.loads(line) + piece = chunk.get("message", {}).get("content") + if piece: + yield piece + if chunk.get("done"): + break + + def available_models(self) -> list[str]: + r = requests.get(f"{self.host}/api/tags", timeout=self.timeout) + r.raise_for_status() + return [m.get("name", "") for m in r.json().get("models", [])] + + +class OllamaEmbedder: + """Local text embeddings via Ollama (default ``nomic-embed-text``), used for + the agent's in-RAM semantic recall. Local + free, so it stays on by default + regardless of which provider answers chat. No key, nothing persisted.""" + + name = "ollama-embed" + + def __init__(self, model: str = "nomic-embed-text", host: str | None = None, + timeout: int = 60, truncate_dim: int | None = 256): + self.model = model + self.host = (host or os.environ.get("OLLAMA_HOST", "http://localhost:11434")).rstrip("/") + self.timeout = timeout + # nomic-embed-text is Matryoshka (MRL)-trained, so its 768-dim vector can + # be truncated to a shorter prefix with little quality loss — faster + # pure-Python cosine and less RAM. Query + stored use the same dim, so + # cosine stays correct. None keeps the full vector. + self.truncate_dim = truncate_dim + + def embed(self, text: str) -> list[float]: + r = requests.post( + f"{self.host}/api/embeddings", + json={"model": self.model, "prompt": text}, + timeout=self.timeout, + ) + r.raise_for_status() + vec = r.json().get("embedding") or [] + if self.truncate_dim is not None: + vec = vec[: self.truncate_dim] + return vec + + +class AnthropicProvider: + """Anthropic Messages API. Cloud — opt-in. Needs ANTHROPIC_API_KEY.""" + + name = "anthropic" + + def __init__(self, model: str = "claude-opus-4-6", api_key: str | None = None, + timeout: int = 120, max_tokens: int = 1024): + self.model = model + self.api_key = api_key or os.environ.get("ANTHROPIC_API_KEY") + self.timeout = timeout + self.max_tokens = max_tokens + if not self.api_key: + raise ValueError("ANTHROPIC_API_KEY not set") + + def complete(self, system: str, messages: list[Msg]) -> str: + payload = { + "model": self.model, + "max_tokens": self.max_tokens, + "system": system, + "messages": [ + {"role": m.role, "content": m.content} + for m in messages + if m.role in ("user", "assistant") + ], + } + r = requests.post( + "https://api.anthropic.com/v1/messages", + json=payload, + timeout=self.timeout, + headers={ + "x-api-key": self.api_key, + "anthropic-version": "2023-06-01", + "content-type": "application/json", + }, + ) + r.raise_for_status() + blocks = r.json().get("content", []) + return "".join(b.get("text", "") for b in blocks).strip() + + def available_models(self) -> list[str]: + r = requests.get( + "https://api.anthropic.com/v1/models", + timeout=self.timeout, + headers={"x-api-key": self.api_key, "anthropic-version": "2023-06-01"}, + ) + r.raise_for_status() + return [m.get("id", "") for m in r.json().get("data", [])] + + +class OpenAICompatibleProvider: + """OpenAI-style /chat/completions — OpenAI, Groq, Together, local vLLM, etc.""" + + name = "openai" + + def __init__(self, model: str = "gpt-4o-mini", api_key: str | None = None, + base_url: str | None = None, timeout: int = 120): + self.model = model + self.api_key = api_key or os.environ.get("OPENAI_API_KEY", "") + self.base_url = (base_url or os.environ.get("OPENAI_BASE_URL", "https://api.openai.com/v1")).rstrip("/") + self.timeout = timeout + + def complete(self, system: str, messages: list[Msg]) -> str: + payload = { + "model": self.model, + "messages": [{"role": "system", "content": system}] + + [{"role": m.role, "content": m.content} for m in messages], + } + headers = {"content-type": "application/json"} + if self.api_key: + headers["authorization"] = f"Bearer {self.api_key}" + r = requests.post( + f"{self.base_url}/chat/completions", json=payload, headers=headers, timeout=self.timeout + ) + r.raise_for_status() + return r.json()["choices"][0]["message"]["content"].strip() + + def available_models(self) -> list[str]: + headers = {} + if self.api_key: + headers["authorization"] = f"Bearer {self.api_key}" + r = requests.get(f"{self.base_url}/models", headers=headers, timeout=self.timeout) + r.raise_for_status() + return [m.get("id", "") for m in r.json().get("data", [])] + + +_BUILTINS = { + "ollama": OllamaProvider, + "anthropic": AnthropicProvider, + "openai": OpenAICompatibleProvider, +} + + +def make_provider(spec: str, model: str | None = None, **opts) -> Provider: + """Build a provider. + + ``spec`` is a builtin name (``ollama`` / ``anthropic`` / ``openai``) or a + ``module:Class`` path to a custom Provider implementation. + """ + if ":" in spec: + mod_name, _, cls_name = spec.partition(":") + cls = getattr(importlib.import_module(mod_name), cls_name) + else: + cls = _BUILTINS.get(spec) + if cls is None: + raise ValueError(f"unknown provider '{spec}' (builtins: {', '.join(_BUILTINS)})") + if model is not None: + opts["model"] = model + return cls(**opts) + + +def preflight(provider: Provider) -> tuple[bool, str]: + """Cheap reachability + model-presence check before joining a room. + + Returns ``(ok, message)``. Lets ``/ai start`` fail fast with a clear reason + (backend down / model not pulled / key missing) instead of erroring on the + first question. Providers without ``available_models`` are assumed reachable. + """ + discover = getattr(provider, "available_models", None) + if discover is None: + return True, f"{provider.name}: no discovery endpoint — assuming reachable" + try: + models = discover() + except Exception as e: # noqa: BLE001 — any failure means "not reachable yet" + return False, f"{provider.name}: cannot reach backend ({e})" + if provider.model in models: + return True, f"{provider.name}/{provider.model}: reachable" + if models: + sample = ", ".join(models[:8]) + more = "…" if len(models) > 8 else "" + return False, ( + f"{provider.name}: model '{provider.model}' not available. " + f"reachable models: {sample}{more}" + ) + return True, f"{provider.name}: reachable (empty model list — skipping check)"