import json from typing import AsyncIterator import httpx from .. import debuglog, endpoints, tlstrust from .base import PromptParts, Provider, ProviderError # Appended after the story text in chat mode, so a chat-tuned model continues # the prose rather than replying conversationally. CHAT_CONTINUE_HINT = "\n\n[Continue the story directly. Output only story text.]" # A machine that is not listening refuses in milliseconds, so a slow connect # means the wrong address rather than a busy model. CONNECT_TIMEOUT = 10.0 # How long to wait for generation when Settings names no value. Upstream # hardcoded 120s, and M1 measured a *cold* load of a 3B model on a GPU-less # four-core host exceeding it three times while the same turn took 6-9 seconds # once the model was resident. 300s covers a cold start on modest hardware and # is still a number: a wedged endpoint fails rather than hanging forever. DEFAULT_READ_TIMEOUT = 300.0 # Embeddings are short and never cold-load a large model. EMBED_READ_TIMEOUT = 60.0 # Completion endpoints have no roles, so a chat has to be flattened into one # labeled transcript that ends on "Assistant:" for the model to continue. _ROLE_LABELS = {"system": "System", "user": "User", "assistant": "Assistant"} def flatten_messages(messages: list[dict]) -> str: turns = "\n\n".join( f"{_ROLE_LABELS.get(m['role'], m['role'])}: {m['content']}" for m in messages ) return f"{turns}\n\nAssistant:" class OpenAICompatibleProvider(Provider): """Adapter for Ollama's OpenAI-compatible `/v1` API. The protocol is OpenAI's, which is what the module is named for; the product speaks it to Ollama and to nothing else. `endpoints.py` decides which addresses may be reached, and every request re-checks — the shape of the wire format is not the same thing as permission to use it. """ def __init__( self, endpoint_url: str, model: str, api_mode: str = "chat", read_timeout: float | None = None, ): self.base_url = endpoint_url.rstrip("/") self.model = model self.api_mode = api_mode # Either "chat" or "completion". # How long to wait for the model, in seconds. Cold-loading a model on a # CPU-only machine can take minutes, and a fixed short timeout reports # that as a failure. See `DEFAULT_READ_TIMEOUT`. self.read_timeout = read_timeout or DEFAULT_READ_TIMEOUT # The token accounting from the last call, when the endpoint reported # any. Every request method writes it, so a caller reads it after the # call it made. One provider is built per request. self.last_usage: dict | None = None def _headers(self) -> dict: # No Authorization header: Ollama does not use one, and this build has # no cloud provider to carry a key for. return {"Content-Type": "application/json"} def _timeout(self, seconds: float | None = None) -> httpx.Timeout: """Short to connect, patient to read. A machine that is not listening says so in milliseconds, so a slow connect is a wrong address rather than a busy model and should fail fast. Generation is the opposite: the first token can be minutes away while a model loads. """ return httpx.Timeout(seconds or self.read_timeout, connect=CONNECT_TIMEOUT) def _record_usage(self, payload: dict) -> None: """Records the endpoint's own token accounting, if it reported any. OpenRouter now always reports usage, and `usage: {include: true}` and `stream_options` are deprecated and do nothing. In a stream the usage arrives on a final chunk that carries no choices, which is why this is read separately from the text extraction. """ usage = payload.get("usage") if isinstance(usage, dict) and usage: self.last_usage = usage def _request(self, parts: PromptParts, temperature: float, max_tokens: int) -> tuple[str, dict]: if self.api_mode == "completion": url = f"{self.base_url}/completions" body = { "model": self.model, "prompt": f"{parts.system}\n\n{parts.story}", "temperature": temperature, "max_tokens": max_tokens, "stream": True, } else: url = f"{self.base_url}/chat/completions" body = { "model": self.model, "messages": [ {"role": "system", "content": parts.system}, {"role": "user", "content": parts.story + CHAT_CONTINUE_HINT}, ], "temperature": temperature, "max_tokens": max_tokens, "stream": True, } return url, body @staticmethod def _extract_chunk(payload: dict) -> str: choices = payload.get("choices") or [] if not choices: return "" choice = choices[0] # A chat stream uses `delta.content`, and a completion stream uses # `text`. The non-stream fallbacks are `message.content` and `text`. delta = choice.get("delta") or {} return ( delta.get("content") or choice.get("text") or (choice.get("message") or {}).get("content") or "" ) @staticmethod def _extract_reasoning(payload: dict) -> str: """Returns a reasoning model's thinking text. OpenRouter normalizes it to `reasoning`, and DeepSeek-style servers use `reasoning_content`. """ choices = payload.get("choices") or [] if not choices: return "" choice = choices[0] delta = choice.get("delta") or {} message = choice.get("message") or {} return ( delta.get("reasoning") or delta.get("reasoning_content") or message.get("reasoning") or message.get("reasoning_content") or "" ) async def generate( self, parts: PromptParts, *, temperature: float, max_tokens: int ) -> AsyncIterator[tuple[str, str]]: """Yields `("text", chunk)` and `("reasoning", chunk)` pairs.""" if not self.model: raise ProviderError("No model configured — set one in Settings.") url, body = self._request(parts, temperature, max_tokens) async for event in self._stream(url, body): yield event async def chat( self, messages: list[dict], *, temperature: float, max_tokens: int ) -> AsyncIterator[tuple[str, str]]: """Runs a plain multi-turn chat, with no story framing and no context assembly. The method sends `[{"role", "content"}, ...]` straight to the endpoint. The AI Chat scratchpad uses it, and the turn engine uses `generate()`. """ if not self.model: raise ProviderError("No model configured — set one in Settings.") if self.api_mode == "completion": url = f"{self.base_url}/completions" body = { "model": self.model, "prompt": flatten_messages(messages), "temperature": temperature, "max_tokens": max_tokens, "stream": True, } else: url = f"{self.base_url}/chat/completions" body = { "model": self.model, "messages": messages, "temperature": temperature, "max_tokens": max_tokens, "stream": True, } async for event in self._stream(url, body): yield event async def _stream(self, url: str, body: dict) -> AsyncIterator[tuple[str, str]]: """Runs the shared SSE request for `generate()` and `chat()`. The method POSTs a streaming request, yields `("text", chunk)` and `("reasoning", chunk)` pairs, and logs the exchange. """ # Re-checked on every request, not only when the endpoint was saved: a # hostname that resolved to a LAN address yesterday can resolve # somewhere else today, and a database row can be edited by hand. reason = endpoints.rejection_reason(url) if reason: raise ProviderError(f"This endpoint can't be used — {reason}.") log = debuglog.start_entry(url, self.model, body) received: list[str] = [] try: async with httpx.AsyncClient( timeout=self._timeout(), verify=tlstrust.ssl_context() ) as client: async with client.stream("POST", url, json=body, headers=self._headers()) as resp: if resp.status_code != 200: detail = (await resp.aread()).decode(errors="replace")[:500] raise ProviderError(self._friendly_http_error(resp.status_code, detail)) # Some servers ignore `stream=true` and return one plain # JSON body, so buffer the non-SSE lines to fall back to. saw_sse = False raw_lines: list[str] = [] async for line in resp.aiter_lines(): if not line.startswith("data:"): if not saw_sse: raw_lines.append(line) continue saw_sse = True data = line[5:].strip() if data == "[DONE]": debuglog.finish_entry( log, response="".join(received), usage=self.last_usage ) return try: payload = json.loads(data) except ValueError: continue self._record_usage(payload) reasoning = self._extract_reasoning(payload) if reasoning: yield "reasoning", reasoning chunk = self._extract_chunk(payload) if chunk: received.append(chunk) yield "text", chunk if not saw_sse: body_text = "\n".join(raw_lines).strip() try: payload = json.loads(body_text) except ValueError: raise ProviderError( "AI endpoint returned neither an SSE stream nor JSON: " + body_text[:200] ) self._record_usage(payload) reasoning = self._extract_reasoning(payload) if reasoning: yield "reasoning", reasoning chunk = self._extract_chunk(payload) if chunk: received.append(chunk) yield "text", chunk if not received: raise ProviderError( "AI endpoint returned a response with no text: " + body_text[:200] ) debuglog.finish_entry(log, response="".join(received), usage=self.last_usage) except httpx.ConnectError as exc: error = f"Could not connect to {self.base_url} — is the AI server running?" debuglog.finish_entry(log, response="".join(received), error=error) raise ProviderError(error) from exc except httpx.TimeoutException as exc: debuglog.finish_entry(log, response="".join(received), error="Timed out") raise ProviderError("The AI endpoint timed out.") from exc except httpx.HTTPError as exc: debuglog.finish_entry(log, response="".join(received), error=str(exc)) raise ProviderError(f"Request to AI endpoint failed: {exc}") from exc except (ProviderError, GeneratorExit, BaseException) as exc: status = "cancelled" if isinstance(exc, GeneratorExit) else str(exc) debuglog.finish_entry(log, response="".join(received), error=status) raise async def complete( self, system: str, user: str, *, temperature: float = 0.3, max_tokens: int = 400 ) -> str: """Runs a single non-streaming completion, for background calls such as summarization. Unlike `generate()`, this adds no story-continuation framing. """ if not self.model: raise ProviderError("No model configured — set one in Settings.") if self.api_mode == "completion": url = f"{self.base_url}/completions" body = { "model": self.model, "prompt": f"{system}\n\n{user}", "temperature": temperature, "max_tokens": max_tokens, "stream": False, } else: url = f"{self.base_url}/chat/completions" body = { "model": self.model, "messages": [ {"role": "system", "content": system}, {"role": "user", "content": user}, ], "temperature": temperature, "max_tokens": max_tokens, "stream": False, } # Same check as `_stream`: every outbound request re-tests the # endpoint, so no path reaches an address the policy refuses. reason = endpoints.rejection_reason(url) if reason: raise ProviderError(f"This endpoint can't be used — {reason}.") log = debuglog.start_entry(url, self.model, body) try: async with httpx.AsyncClient( timeout=self._timeout(), verify=tlstrust.ssl_context() ) as client: resp = await client.post(url, json=body, headers=self._headers()) except httpx.HTTPError as exc: debuglog.finish_entry(log, error=str(exc)) raise ProviderError(f"Request to AI endpoint failed: {exc}") from exc if resp.status_code != 200: error = self._friendly_http_error(resp.status_code, resp.text[:500]) debuglog.finish_entry(log, error=error) raise ProviderError(error) try: payload = resp.json() except ValueError as exc: debuglog.finish_entry(log, error="Invalid JSON response") raise ProviderError("AI endpoint returned invalid JSON.") from exc self._record_usage(payload) text = self._extract_chunk(payload) debuglog.finish_entry(log, response=text, usage=self.last_usage) return text.strip() async def embed(self, texts: list[str]) -> list[list[float]]: """POSTs to /v1/embeddings. Here `self.model` is the embedding model.""" if not self.model: raise ProviderError("No embedding model configured — set one in Settings.") url = f"{self.base_url}/embeddings" body = {"model": self.model, "input": texts} # Same check as `_stream`: every outbound request re-tests the # endpoint, so no path reaches an address the policy refuses. reason = endpoints.rejection_reason(url) if reason: raise ProviderError(f"This endpoint can't be used — {reason}.") log = debuglog.start_entry(url, self.model, body) try: async with httpx.AsyncClient( timeout=self._timeout(EMBED_READ_TIMEOUT), verify=tlstrust.ssl_context() ) as client: resp = await client.post(url, json=body, headers=self._headers()) except httpx.HTTPError as exc: debuglog.finish_entry(log, error=str(exc)) raise ProviderError(f"Embedding request failed: {exc}") from exc if resp.status_code != 200: error = self._friendly_http_error(resp.status_code, resp.text[:500]) debuglog.finish_entry(log, error=error) raise ProviderError(error) try: data = resp.json().get("data", []) vectors = [item["embedding"] for item in sorted(data, key=lambda d: d.get("index", 0))] except (ValueError, KeyError, TypeError) as exc: debuglog.finish_entry(log, error="Malformed embeddings response") raise ProviderError("AI endpoint returned malformed embeddings.") from exc if len(vectors) != len(texts): debuglog.finish_entry(log, error="Embedding count mismatch") raise ProviderError("AI endpoint returned the wrong number of embeddings.") debuglog.finish_entry(log, response=f"{len(vectors)} vectors × {len(vectors[0]) if vectors else 0} dims") return vectors def _friendly_http_error(self, status: int, detail: str) -> str: if status == 401: return "Authentication failed — check your API key in Settings." if status == 404: return ( f"Endpoint or model not found (HTTP 404). Check the endpoint URL and that " f"model '{self.model}' exists. {detail}" ) if status == 429: # OpenRouter's shared free tier has a per-day cap. Distinguish it # from a short-term burst limit, so the message tells the reader what # to do. if "free-models-per-day" in detail: return ( "The free demo has hit its daily request limit (resets at " "00:00 UTC). Please try again later." ) return "The AI is getting too many requests right now — wait a moment and try again." return f"AI endpoint returned HTTP {status}: {detail}"