Initial commit: AI Dungeon clone (FastAPI backend + React frontend)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01KFsGHju9szibJJa2YJcdbg
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
from .base import PromptParts, Provider, ProviderError
|
||||
from .openai_compatible import OpenAICompatibleProvider
|
||||
|
||||
__all__ = ["PromptParts", "Provider", "ProviderError", "OpenAICompatibleProvider"]
|
||||
@@ -0,0 +1,27 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import AsyncIterator
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptParts:
|
||||
"""Assembled context, provider-agnostic. Providers map this to their wire format."""
|
||||
|
||||
system: str # narrator prompt + AI instructions + memory
|
||||
story: str # the story text so far (already token-budgeted)
|
||||
|
||||
|
||||
class ProviderError(Exception):
|
||||
"""User-presentable provider failure (connection refused, bad key, model not found…)."""
|
||||
|
||||
|
||||
class Provider(ABC):
|
||||
@abstractmethod
|
||||
def generate(
|
||||
self,
|
||||
parts: PromptParts,
|
||||
*,
|
||||
temperature: float,
|
||||
max_tokens: int,
|
||||
) -> AsyncIterator[tuple[str, str]]:
|
||||
"""Yield ("text" | "reasoning", chunk) pairs. Raises ProviderError on failure."""
|
||||
@@ -0,0 +1,242 @@
|
||||
import json
|
||||
from typing import AsyncIterator
|
||||
|
||||
import httpx
|
||||
|
||||
from .. import debuglog
|
||||
from .base import PromptParts, Provider, ProviderError
|
||||
|
||||
# Framing appended after the story text in chat mode, so chat-tuned models keep
|
||||
# continuing prose instead of replying conversationally.
|
||||
CHAT_CONTINUE_HINT = "\n\n[Continue the story directly. Output only story text.]"
|
||||
|
||||
|
||||
class OpenAICompatibleProvider(Provider):
|
||||
"""Adapter for any /v1-style endpoint: Ollama, LM Studio, OpenAI, OpenRouter, vLLM, Groq…"""
|
||||
|
||||
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 # "chat" | "completion"
|
||||
# Thinking budget for reasoning models, on top of max_tokens. 0 = the
|
||||
# `reasoning` param is not sent (endpoints that don't know it may
|
||||
# reject unknown fields).
|
||||
self.reasoning_max_tokens = reasoning_max_tokens
|
||||
|
||||
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:
|
||||
"""Give reasoning models their own thinking budget (OpenRouter-style),
|
||||
raising max_tokens so the actual output keeps its full budget."""
|
||||
if self.reasoning_max_tokens > 0 and self.api_mode == "chat":
|
||||
body["reasoning"] = {"max_tokens": self.reasoning_max_tokens}
|
||||
body["max_tokens"] += self.reasoning_max_tokens
|
||||
|
||||
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)
|
||||
return url, body
|
||||
|
||||
@staticmethod
|
||||
def _extract_chunk(payload: dict) -> str:
|
||||
choices = payload.get("choices") or []
|
||||
if not choices:
|
||||
return ""
|
||||
choice = choices[0]
|
||||
# chat stream → delta.content; completion stream → text;
|
||||
# non-stream fallbacks → message.content / 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:
|
||||
"""Reasoning-model thinking: OpenRouter normalizes to `reasoning`;
|
||||
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" | "reasoning", chunk) pairs."""
|
||||
if not self.model:
|
||||
raise ProviderError("No model configured — set one in Settings.")
|
||||
url, body = self._request(parts, temperature, max_tokens)
|
||||
|
||||
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))
|
||||
async for line in resp.aiter_lines():
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
data = line[5:].strip()
|
||||
if data == "[DONE]":
|
||||
debuglog.finish_entry(log, response="".join(received))
|
||||
return
|
||||
try:
|
||||
payload = json.loads(data)
|
||||
except ValueError:
|
||||
continue
|
||||
reasoning = self._extract_reasoning(payload)
|
||||
if reasoning:
|
||||
yield "reasoning", reasoning
|
||||
chunk = self._extract_chunk(payload)
|
||||
if chunk:
|
||||
received.append(chunk)
|
||||
yield "text", chunk
|
||||
debuglog.finish_entry(log, response="".join(received))
|
||||
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:
|
||||
"""Single non-streaming completion for background calls (summarization).
|
||||
Unlike generate(), no story-continuation framing is added."""
|
||||
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)
|
||||
|
||||
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:
|
||||
text = self._extract_chunk(resp.json())
|
||||
except ValueError as exc:
|
||||
debuglog.finish_entry(log, error="Invalid JSON response")
|
||||
raise ProviderError("AI endpoint returned invalid JSON.") from exc
|
||||
debuglog.finish_entry(log, response=text)
|
||||
return text.strip()
|
||||
|
||||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||
"""POST /v1/embeddings; self.model is the embedding model here."""
|
||||
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}"
|
||||
)
|
||||
return f"AI endpoint returned HTTP {status}: {detail}"
|
||||
Reference in New Issue
Block a user