Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c2aceffda1 | |||
| e627356168 | |||
| 67517926aa | |||
| 01e6975d20 | |||
| 4a60171035 |
@@ -2,3 +2,9 @@
|
|||||||
|
|
||||||
# v1 runtime data (DB, backups, snapshots, exports, config with secrets)
|
# v1 runtime data (DB, backups, snapshots, exports, config with secrets)
|
||||||
data/
|
data/
|
||||||
|
|
||||||
|
# Python
|
||||||
|
.venv/
|
||||||
|
__pycache__/
|
||||||
|
*.pyc
|
||||||
|
.pytest_cache/
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
3.12
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
from fastapi import FastAPI
|
||||||
|
|
||||||
|
app = FastAPI(title="chat")
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/health")
|
||||||
|
def health():
|
||||||
|
return {"status": "ok"}
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
import os
|
||||||
|
import tomllib
|
||||||
|
from pathlib import Path
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
REPO_ROOT = Path(__file__).resolve().parent.parent
|
||||||
|
DEFAULT_CONFIG = REPO_ROOT / "data" / "config.toml"
|
||||||
|
DEFAULT_DB = REPO_ROOT / "data" / "chat.db"
|
||||||
|
|
||||||
|
class Settings(BaseModel):
|
||||||
|
featherless_api_key: str
|
||||||
|
featherless_base_url: str = "https://api.featherless.ai/v1"
|
||||||
|
narrative_model: str = "dphn/Dolphin-Mistral-24B-Venice-Edition"
|
||||||
|
classifier_model: str = "NousResearch/Hermes-3-Llama-3.1-8B"
|
||||||
|
classifier_fallbacks: list[str] = Field(
|
||||||
|
default_factory=lambda: [
|
||||||
|
"cognitivecomputations/dolphin-2.9.4-llama3-8b",
|
||||||
|
"mlabonne/Meta-Llama-3.1-8B-Instruct-abliterated",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
ooc_marker: str = "(("
|
||||||
|
retrieval_k: int = 4
|
||||||
|
narrative_budget_hard: int = 8000
|
||||||
|
narrative_budget_soft: int = 6000
|
||||||
|
classifier_budget_hard: int = 4000
|
||||||
|
classifier_timeout_s: float = 10.0
|
||||||
|
db_path: Path = DEFAULT_DB
|
||||||
|
data_dir: Path = REPO_ROOT / "data"
|
||||||
|
bind_host: str = "127.0.0.1"
|
||||||
|
bind_port: int = 8000
|
||||||
|
|
||||||
|
def load_settings() -> Settings:
|
||||||
|
config_path = Path(os.environ.get("CHAT_CONFIG_PATH", DEFAULT_CONFIG))
|
||||||
|
raw: dict = {}
|
||||||
|
if config_path.exists():
|
||||||
|
raw = tomllib.loads(config_path.read_text())
|
||||||
|
if "CHAT_DB_PATH" in os.environ:
|
||||||
|
raw["db_path"] = Path(os.environ["CHAT_DB_PATH"])
|
||||||
|
return Settings(**raw)
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
import sqlite3
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def open_db(path: Path):
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
conn = sqlite3.connect(path)
|
||||||
|
conn.execute("PRAGMA journal_mode=WAL")
|
||||||
|
conn.execute("PRAGMA foreign_keys=ON")
|
||||||
|
try:
|
||||||
|
yield conn
|
||||||
|
conn.commit()
|
||||||
|
finally:
|
||||||
|
conn.close()
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from chat.db.connection import open_db
|
||||||
|
|
||||||
|
MIGRATIONS_DIR = Path(__file__).parent / "migrations"
|
||||||
|
|
||||||
|
|
||||||
|
def apply_migrations(db_path: Path) -> None:
|
||||||
|
with open_db(db_path) as conn:
|
||||||
|
conn.execute(
|
||||||
|
"CREATE TABLE IF NOT EXISTS meta (key TEXT PRIMARY KEY, value TEXT)"
|
||||||
|
)
|
||||||
|
cur = conn.execute("SELECT value FROM meta WHERE key = 'schema_version'")
|
||||||
|
row = cur.fetchone()
|
||||||
|
current = int(row[0]) if row else 0
|
||||||
|
for path in sorted(MIGRATIONS_DIR.glob("*.sql")):
|
||||||
|
version = int(path.stem.split("_", 1)[0])
|
||||||
|
if version <= current:
|
||||||
|
continue
|
||||||
|
sql = path.read_text()
|
||||||
|
conn.executescript(sql)
|
||||||
|
conn.execute(
|
||||||
|
"INSERT OR REPLACE INTO meta (key, value) VALUES ('schema_version', ?)",
|
||||||
|
(str(version),),
|
||||||
|
)
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
-- meta table is created by the migrate runner; this migration is a marker.
|
||||||
|
SELECT 1;
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
CREATE TABLE classifier_failures (
|
||||||
|
id INTEGER PRIMARY KEY,
|
||||||
|
kind TEXT NOT NULL,
|
||||||
|
model TEXT NOT NULL,
|
||||||
|
raw_text TEXT,
|
||||||
|
attempt_count INTEGER NOT NULL,
|
||||||
|
created_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||||
|
);
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
import json
|
||||||
|
import asyncio
|
||||||
|
from typing import TypeVar
|
||||||
|
from pydantic import BaseModel, ValidationError
|
||||||
|
from .client import LLMClient, Message
|
||||||
|
|
||||||
|
T = TypeVar("T", bound=BaseModel)
|
||||||
|
|
||||||
|
REFUSAL_PATTERNS = ("i can't", "i cannot", "i'm sorry, but", "as an ai")
|
||||||
|
|
||||||
|
|
||||||
|
async def classify(
|
||||||
|
client: LLMClient,
|
||||||
|
*,
|
||||||
|
model: str,
|
||||||
|
system: str,
|
||||||
|
user: str,
|
||||||
|
schema: type[T],
|
||||||
|
default: T | None = None,
|
||||||
|
timeout_s: float = 10.0,
|
||||||
|
) -> T:
|
||||||
|
msgs = [
|
||||||
|
Message(role="system", content=system + "\n\nRespond with JSON only matching the schema."),
|
||||||
|
Message(role="user", content=user),
|
||||||
|
]
|
||||||
|
for attempt in range(2):
|
||||||
|
try:
|
||||||
|
text = await asyncio.wait_for(
|
||||||
|
client.generate(msgs, model=model, response_format={"type": "json_object"}),
|
||||||
|
timeout=timeout_s,
|
||||||
|
)
|
||||||
|
if any(p in text.lower()[:80] for p in REFUSAL_PATTERNS) and not text.strip().startswith("{"):
|
||||||
|
raise ValueError("refusal-shaped response")
|
||||||
|
return schema.model_validate_json(text)
|
||||||
|
except (ValidationError, ValueError, json.JSONDecodeError, asyncio.TimeoutError):
|
||||||
|
msgs[0] = Message(role="system", content=system + "\n\nRespond with valid JSON ONLY. No prose.")
|
||||||
|
continue
|
||||||
|
if default is None:
|
||||||
|
raise RuntimeError(f"classify failed for schema {schema.__name__} with no default")
|
||||||
|
return default
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Protocol, AsyncIterator, Sequence
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Message:
|
||||||
|
role: str # "system" | "user" | "assistant"
|
||||||
|
content: str
|
||||||
|
|
||||||
|
|
||||||
|
class LLMClient(Protocol):
|
||||||
|
async def generate(self, messages: Sequence[Message], *, model: str, **params) -> str: ...
|
||||||
|
def stream(self, messages: Sequence[Message], *, model: str, **params) -> AsyncIterator[str]: ...
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from typing import AsyncIterator, Sequence
|
||||||
|
from openai import AsyncOpenAI
|
||||||
|
from .client import Message
|
||||||
|
|
||||||
|
|
||||||
|
class FeatherlessClient:
|
||||||
|
def __init__(self, api_key: str, base_url: str = "https://api.featherless.ai/v1"):
|
||||||
|
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url)
|
||||||
|
|
||||||
|
async def generate(self, messages: Sequence[Message], *, model: str, **params) -> str:
|
||||||
|
resp = await self._client.chat.completions.create(
|
||||||
|
model=model,
|
||||||
|
messages=[{"role": m.role, "content": m.content} for m in messages],
|
||||||
|
**params,
|
||||||
|
)
|
||||||
|
return resp.choices[0].message.content or ""
|
||||||
|
|
||||||
|
async def stream(self, messages: Sequence[Message], *, model: str, **params) -> AsyncIterator[str]:
|
||||||
|
stream = await self._client.chat.completions.create(
|
||||||
|
model=model,
|
||||||
|
messages=[{"role": m.role, "content": m.content} for m in messages],
|
||||||
|
stream=True,
|
||||||
|
**params,
|
||||||
|
)
|
||||||
|
async for chunk in stream:
|
||||||
|
delta = chunk.choices[0].delta.content or ""
|
||||||
|
if delta:
|
||||||
|
yield delta
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from typing import AsyncIterator, Sequence
|
||||||
|
from .client import Message
|
||||||
|
|
||||||
|
|
||||||
|
class MockLLMClient:
|
||||||
|
def __init__(self, canned: list[str]):
|
||||||
|
self._canned = list(canned)
|
||||||
|
|
||||||
|
async def generate(self, messages: Sequence[Message], *, model: str, **params) -> str:
|
||||||
|
return self._canned.pop(0)
|
||||||
|
|
||||||
|
async def stream(self, messages: Sequence[Message], *, model: str, **params) -> AsyncIterator[str]:
|
||||||
|
text = self._canned.pop(0)
|
||||||
|
for ch in text:
|
||||||
|
yield ch
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
# Copy this file to data/config.toml and fill in your API key.
|
||||||
|
featherless_api_key = "REPLACE_ME"
|
||||||
|
narrative_model = "dphn/Dolphin-Mistral-24B-Venice-Edition"
|
||||||
|
classifier_model = "NousResearch/Hermes-3-Llama-3.1-8B"
|
||||||
|
ooc_marker = "(("
|
||||||
|
retrieval_k = 4
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
[project]
|
||||||
|
name = "chat"
|
||||||
|
version = "0.1.0"
|
||||||
|
requires-python = ">=3.11"
|
||||||
|
dependencies = [
|
||||||
|
"fastapi>=0.110",
|
||||||
|
"uvicorn[standard]>=0.30",
|
||||||
|
"httpx>=0.27",
|
||||||
|
"pydantic>=2.6",
|
||||||
|
"pydantic-settings>=2.2",
|
||||||
|
"openai>=1.30",
|
||||||
|
"instructor>=1.3",
|
||||||
|
"tiktoken>=0.7",
|
||||||
|
"jinja2>=3.1",
|
||||||
|
"aiosqlite>=0.20",
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = ["pytest>=8", "pytest-asyncio>=0.23", "freezegun>=1.4"]
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
pythonpath = ["."]
|
||||||
|
asyncio_mode = "auto"
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
import pytest
|
||||||
|
from pydantic import BaseModel
|
||||||
|
from chat.llm.mock import MockLLMClient
|
||||||
|
from chat.llm.classify import classify
|
||||||
|
|
||||||
|
|
||||||
|
class Verdict(BaseModel):
|
||||||
|
score: int
|
||||||
|
reason: str
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_classify_parses_valid_json():
|
||||||
|
mock = MockLLMClient(canned=['{"score": 2, "reason": "notable"}'])
|
||||||
|
result = await classify(mock, model="m", system="x", user="y", schema=Verdict)
|
||||||
|
assert result.score == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_classify_falls_back_on_unparseable_after_retry():
|
||||||
|
mock = MockLLMClient(canned=["nope", "still nope"])
|
||||||
|
default = Verdict(score=1, reason="fallback")
|
||||||
|
result = await classify(mock, model="m", system="x", user="y", schema=Verdict, default=default)
|
||||||
|
assert result.reason == "fallback"
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
import pytest
|
||||||
|
from chat.config import load_settings
|
||||||
|
|
||||||
|
def test_load_settings_reads_toml(tmp_path, monkeypatch):
|
||||||
|
cfg = tmp_path / "config.toml"
|
||||||
|
cfg.write_text("""
|
||||||
|
featherless_api_key = "sk-test"
|
||||||
|
narrative_model = "dphn/Dolphin-Mistral-24B-Venice-Edition"
|
||||||
|
classifier_model = "NousResearch/Hermes-3-Llama-3.1-8B"
|
||||||
|
ooc_marker = "(("
|
||||||
|
retrieval_k = 4
|
||||||
|
""")
|
||||||
|
monkeypatch.setenv("CHAT_CONFIG_PATH", str(cfg))
|
||||||
|
s = load_settings()
|
||||||
|
assert s.featherless_api_key == "sk-test"
|
||||||
|
assert s.narrative_model.startswith("dphn/")
|
||||||
|
assert s.retrieval_k == 4
|
||||||
|
|
||||||
|
def test_chat_db_path_env_overrides_default(tmp_path, monkeypatch):
|
||||||
|
monkeypatch.setenv("CHAT_DB_PATH", str(tmp_path / "alt.db"))
|
||||||
|
monkeypatch.setenv("CHAT_CONFIG_PATH", str(tmp_path / "config.toml"))
|
||||||
|
(tmp_path / "config.toml").write_text('featherless_api_key = "x"\n')
|
||||||
|
s = load_settings()
|
||||||
|
assert s.db_path == tmp_path / "alt.db"
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
from fastapi.testclient import TestClient
|
||||||
|
from chat.app import app
|
||||||
|
|
||||||
|
|
||||||
|
def test_health_endpoint_returns_ok():
|
||||||
|
client = TestClient(app)
|
||||||
|
response = client.get("/health")
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {"status": "ok"}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
import pytest
|
||||||
|
from chat.llm.mock import MockLLMClient
|
||||||
|
from chat.llm.client import Message
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mock_returns_canned_response():
|
||||||
|
client = MockLLMClient(canned=["Hello, world."])
|
||||||
|
msgs = [Message(role="user", content="hi")]
|
||||||
|
out = await client.generate(msgs, model="any")
|
||||||
|
assert out == "Hello, world."
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mock_streams_tokens():
|
||||||
|
client = MockLLMClient(canned=["abcd"])
|
||||||
|
msgs = [Message(role="user", content="hi")]
|
||||||
|
chunks = []
|
||||||
|
async for chunk in client.stream(msgs, model="any"):
|
||||||
|
chunks.append(chunk)
|
||||||
|
assert "".join(chunks) == "abcd"
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
from chat.db.connection import open_db
|
||||||
|
from chat.db.migrate import apply_migrations
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_migrations_creates_meta_table(tmp_path):
|
||||||
|
db = tmp_path / "test.db"
|
||||||
|
apply_migrations(db)
|
||||||
|
with open_db(db) as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT value FROM meta WHERE key = 'schema_version'"
|
||||||
|
).fetchone()
|
||||||
|
assert row is not None
|
||||||
|
assert int(row[0]) >= 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_migrations_idempotent(tmp_path):
|
||||||
|
db = tmp_path / "test.db"
|
||||||
|
apply_migrations(db)
|
||||||
|
apply_migrations(db) # second call must be a no-op
|
||||||
|
with open_db(db) as conn:
|
||||||
|
count = conn.execute("SELECT COUNT(*) FROM meta").fetchone()[0]
|
||||||
|
assert count == 1
|
||||||
Reference in New Issue
Block a user