* Rewrite comments in Google developer documentation style Rewrite the comments and docstrings across the backend core modules so they read plainly. The previous prose was accurate but dense and figurative, which made it slow to skim. Applies the Google developer documentation style guide: short sentences, active voice, present tense, American spelling, and no metaphors, idioms, or rhetorical asides. Replaces em-dash chains with separate sentences.
424 lines
19 KiB
Python
424 lines
19 KiB
Python
import json
|
||
from typing import AsyncIterator
|
||
|
||
import httpx
|
||
|
||
from .. import debuglog, netguard
|
||
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.]"
|
||
|
||
# OpenRouter serves one model from whichever upstream is available, and every
|
||
# upstream holds its own prompt cache, so a request routed somewhere new starts
|
||
# with a cold cache however stable the prompt is. Naming a preferred upstream
|
||
# makes routing deterministic, which is what allows a cache hit at all.
|
||
#
|
||
# `allow_fallbacks` stays at its default of true on purpose, because this is a
|
||
# preference rather than a restriction. If the named upstream is down, the
|
||
# request still goes elsewhere and only misses the cache, which is the behavior
|
||
# without this setting.
|
||
#
|
||
# This is a list rather than a value derived from the model slug. The vendor half
|
||
# of a slug is usually the provider slug, such as "deepseek/..." mapping to
|
||
# "deepseek", which was verified against /api/v1/providers, but not reliably.
|
||
# Google's models are served by "google-ai-studio" and "google-vertex", and there
|
||
# is no "google". Look a vendor up on the model's Providers tab before adding it
|
||
# here. A slug that does not exist is a routing preference that, at best, does
|
||
# nothing.
|
||
_OPENROUTER_HOST = "openrouter.ai"
|
||
_PREFERRED_UPSTREAM = {"deepseek": "deepseek"}
|
||
|
||
|
||
# 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 any /v1-style endpoint.
|
||
|
||
This covers Ollama, LM Studio, OpenAI, OpenRouter, vLLM, and Groq, among
|
||
others.
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
endpoint_url: str,
|
||
api_key: str,
|
||
model: str,
|
||
api_mode: str = "chat",
|
||
reasoning_max_tokens: int = 0,
|
||
):
|
||
self.base_url = endpoint_url.rstrip("/")
|
||
self.api_key = api_key
|
||
self.model = model
|
||
self.api_mode = api_mode # Either "chat" or "completion".
|
||
# The thinking budget for reasoning models, on top of `max_tokens`. A
|
||
# value of 0 means the `reasoning` parameter is not sent, because an
|
||
# endpoint that does not know the field may reject it. A negative value
|
||
# asks the endpoint to turn reasoning off.
|
||
self.reasoning_max_tokens = reasoning_max_tokens
|
||
# The token accounting from the last call, when the endpoint reported
|
||
# any. It holds the prompt and completion counts, plus, on OpenRouter,
|
||
# `prompt_tokens_details.cached_tokens`, which is the number of prompt
|
||
# tokens read from cache rather than billed in full. 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:
|
||
headers = {"Content-Type": "application/json"}
|
||
if self.api_key:
|
||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||
return headers
|
||
|
||
def _apply_reasoning_budget(self, body: dict) -> None:
|
||
"""Gives reasoning models their own thinking budget, in the OpenRouter style.
|
||
|
||
The method raises `max_tokens`, so the output keeps its full budget.
|
||
|
||
A negative budget does the opposite. It sends `effort: "none"` to turn
|
||
reasoning off on a model that reasons by default, such as DeepSeek V4
|
||
Flash. That differs from `exclude: true`, which still reasons and still
|
||
bills for it while hiding the trace. Zero still means send nothing, so an
|
||
endpoint that rejects unknown fields, such as Ollama, keeps working.
|
||
"""
|
||
if self.api_mode != "chat":
|
||
return
|
||
if self.reasoning_max_tokens < 0:
|
||
body["reasoning"] = {"effort": "none"}
|
||
elif self.reasoning_max_tokens > 0:
|
||
body["reasoning"] = {"max_tokens": self.reasoning_max_tokens}
|
||
body["max_tokens"] += self.reasoning_max_tokens
|
||
|
||
def _apply_provider_routing(self, body: dict) -> None:
|
||
"""Prefers one upstream on OpenRouter, so the prompt cache stays warm.
|
||
|
||
The method does nothing anywhere else. `provider` is an OpenRouter
|
||
extension, and Ollama and similar servers reject fields they do not know.
|
||
The `reasoning` parameter above is written around the same constraint.
|
||
"""
|
||
if _OPENROUTER_HOST not in self.base_url:
|
||
return
|
||
upstream = _PREFERRED_UPSTREAM.get(self.model.split("/", 1)[0].lower())
|
||
if upstream:
|
||
body["provider"] = {"order": [upstream]}
|
||
|
||
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,
|
||
}
|
||
self._apply_reasoning_budget(body)
|
||
self._apply_provider_routing(body)
|
||
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,
|
||
}
|
||
self._apply_reasoning_budget(body)
|
||
self._apply_provider_routing(body)
|
||
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.
|
||
"""
|
||
# SSRF guard for hosted mode. A user-supplied `endpoint_url` must not
|
||
# point at an internal or metadata address. This does nothing for a
|
||
# local install.
|
||
reason = netguard.endpoint_block_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=httpx.Timeout(120, connect=10)) 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,
|
||
}
|
||
self._apply_reasoning_budget(body)
|
||
self._apply_provider_routing(body)
|
||
|
||
log = debuglog.start_entry(url, self.model, body)
|
||
try:
|
||
async with httpx.AsyncClient(timeout=httpx.Timeout(120, connect=10)) 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}
|
||
log = debuglog.start_entry(url, self.model, body)
|
||
try:
|
||
async with httpx.AsyncClient(timeout=httpx.Timeout(60, connect=10)) 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}"
|