Big poker mode changes and hotfixes. #5
+29
-14
@@ -37,30 +37,45 @@ def _resolved_model(cfg, backend: Backend, model: str | None) -> str:
|
|||||||
return model or cfg.local_model
|
return model or cfg.local_model
|
||||||
|
|
||||||
|
|
||||||
def complete(messages: list[Message], backend: Backend = "local", model: str | None = None) -> str:
|
def complete(messages: list[Message], backend: Backend = "local", model: str | None = None,
|
||||||
|
max_tokens: int | None = None, timeout: float | None = None) -> str:
|
||||||
"""Generate a completion. `model` overrides the backend's default model
|
"""Generate a completion. `model` overrides the backend's default model
|
||||||
(used so live chat can run a stronger cloud model than bulk consolidation)."""
|
(used so live chat can run a stronger cloud model than bulk consolidation).
|
||||||
|
|
||||||
|
`max_tokens` caps the generation length (guards a slow local model against
|
||||||
|
rambling for thousands of tokens). `timeout`, when set, bounds each request
|
||||||
|
and disables the SDK's own retries so the caller owns retry/fallback policy.
|
||||||
|
Both default to None → unchanged behavior for every existing caller."""
|
||||||
cfg = load()
|
cfg = load()
|
||||||
mdl = _resolved_model(cfg, backend, model)
|
mdl = _resolved_model(cfg, backend, model)
|
||||||
logbus.log("info", "llm call", kind="complete", backend=backend, model=mdl, tok=_approx_tok(messages))
|
logbus.log("info", "llm call", kind="complete", backend=backend, model=mdl, tok=_approx_tok(messages))
|
||||||
t0 = time.monotonic()
|
t0 = time.monotonic()
|
||||||
|
|
||||||
if backend == "cloud":
|
if backend in ("cloud", "mi50"):
|
||||||
if not cfg.openai_api_key:
|
if backend == "cloud":
|
||||||
raise RuntimeError("OPENAI_API_KEY is not set")
|
if not cfg.openai_api_key:
|
||||||
client = OpenAI(api_key=cfg.openai_api_key)
|
raise RuntimeError("OPENAI_API_KEY is not set")
|
||||||
resp = client.chat.completions.create(model=mdl, messages=messages)
|
client_kwargs: dict = {"api_key": cfg.openai_api_key}
|
||||||
out = resp.choices[0].message.content or ""
|
else:
|
||||||
elif backend == "mi50":
|
# MI50 box runs an OpenAI-compatible llama.cpp server; key is unused.
|
||||||
# MI50 box runs an OpenAI-compatible llama.cpp server; key is unused.
|
client_kwargs = {"api_key": "not-needed", "base_url": cfg.mi50_base_url}
|
||||||
client = OpenAI(api_key="not-needed", base_url=cfg.mi50_base_url)
|
if timeout is not None:
|
||||||
resp = client.chat.completions.create(model=mdl, messages=messages)
|
client_kwargs["timeout"] = timeout
|
||||||
|
client_kwargs["max_retries"] = 0 # caller owns retries (see summary.py)
|
||||||
|
client = OpenAI(**client_kwargs)
|
||||||
|
create_kwargs: dict = {"model": mdl, "messages": messages}
|
||||||
|
if max_tokens is not None:
|
||||||
|
create_kwargs["max_tokens"] = max_tokens
|
||||||
|
resp = client.chat.completions.create(**create_kwargs)
|
||||||
out = resp.choices[0].message.content or ""
|
out = resp.choices[0].message.content or ""
|
||||||
else:
|
else:
|
||||||
|
payload: dict = {"model": mdl, "messages": messages, "stream": False}
|
||||||
|
if max_tokens is not None:
|
||||||
|
payload["options"] = {"num_predict": max_tokens}
|
||||||
resp = httpx.post(
|
resp = httpx.post(
|
||||||
f"{cfg.local_base_url}/api/chat",
|
f"{cfg.local_base_url}/api/chat",
|
||||||
json={"model": mdl, "messages": messages, "stream": False},
|
json=payload,
|
||||||
timeout=120,
|
timeout=timeout or 120,
|
||||||
)
|
)
|
||||||
resp.raise_for_status()
|
resp.raise_for_status()
|
||||||
out = resp.json()["message"]["content"]
|
out = resp.json()["message"]["content"]
|
||||||
|
|||||||
+34
-9
@@ -17,7 +17,16 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
|
|||||||
from lyra import config, llm, logbus, memory
|
from lyra import config, llm, logbus, memory
|
||||||
from lyra.llm import Backend, Message
|
from lyra.llm import Backend, Message
|
||||||
|
|
||||||
_RETRIES = 4
|
# Consolidation LLM budget. A gist is short (a handful of sentences), so cap the
|
||||||
|
# generation hard — an uncapped local model will otherwise ramble for thousands
|
||||||
|
# of tokens and, on a slow GPU, blow the request timeout. 768 is ~3x the longest
|
||||||
|
# real gist we've stored.
|
||||||
|
SUMMARY_MAX_TOKENS = 768
|
||||||
|
# Attempts on the primary backend before falling back to cloud.
|
||||||
|
MI50_ATTEMPTS = 2
|
||||||
|
# Per-call timeout (seconds). A capped 768-token gist finishes in ~60-90s on the
|
||||||
|
# MI50; 150s is headroom but bails a hung call fast so fallback isn't slow.
|
||||||
|
SUMMARY_TIMEOUT = 150
|
||||||
|
|
||||||
# Re-summarize a session once it has accumulated this many new raw exchanges.
|
# Re-summarize a session once it has accumulated this many new raw exchanges.
|
||||||
SUMMARIZE_AFTER = 20
|
SUMMARIZE_AFTER = 20
|
||||||
@@ -61,16 +70,32 @@ def _summarize_text(text: str, backend: Backend) -> str:
|
|||||||
{"role": "system", "content": _PROMPT},
|
{"role": "system", "content": _PROMPT},
|
||||||
{"role": "user", "content": text},
|
{"role": "user", "content": text},
|
||||||
]
|
]
|
||||||
# Retry transient backend errors (e.g. the GPU server restarting) with backoff.
|
|
||||||
for attempt in range(_RETRIES):
|
def _call(be: Backend) -> str:
|
||||||
|
return llm.complete(messages, backend=be,
|
||||||
|
max_tokens=SUMMARY_MAX_TOKENS, timeout=SUMMARY_TIMEOUT)
|
||||||
|
|
||||||
|
# Try the primary backend a bounded number of times (each call fast-fails via
|
||||||
|
# SUMMARY_TIMEOUT), with a short backoff for a transient blip / restarting GPU.
|
||||||
|
last_exc: Exception | None = None
|
||||||
|
for attempt in range(MI50_ATTEMPTS):
|
||||||
try:
|
try:
|
||||||
return llm.complete(messages, backend=backend)
|
return _call(backend)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
if attempt == _RETRIES - 1:
|
last_exc = exc
|
||||||
raise
|
logbus.log("debug", "summary retry", attempt=attempt + 1,
|
||||||
logbus.log("debug", "summary retry", attempt=attempt + 1, error=str(exc)[:80])
|
backend=backend, error=str(exc)[:80])
|
||||||
time.sleep(5 * (attempt + 1))
|
if attempt < MI50_ATTEMPTS - 1:
|
||||||
raise RuntimeError("unreachable")
|
time.sleep(5 * (attempt + 1))
|
||||||
|
|
||||||
|
# Primary exhausted. If it wasn't already cloud and cloud is configured, fall
|
||||||
|
# back once so a stuck/offline MI50 doesn't sink consolidation for the night.
|
||||||
|
if backend != "cloud" and config.load().openai_api_key:
|
||||||
|
logbus.log("info", "summary fell back to cloud", primary=backend,
|
||||||
|
error=str(last_exc)[:80] if last_exc else None)
|
||||||
|
return _call("cloud")
|
||||||
|
|
||||||
|
raise last_exc if last_exc else RuntimeError("summary failed")
|
||||||
|
|
||||||
|
|
||||||
def _summarize_transcript(transcript: str, backend: Backend) -> str:
|
def _summarize_transcript(transcript: str, backend: Backend) -> str:
|
||||||
|
|||||||
+1
-1
@@ -20,7 +20,7 @@ def lyra(tmp_path, monkeypatch):
|
|||||||
# reflect() expects JSON back; everything else just stores the text.
|
# reflect() expects JSON back; everything else just stores the text.
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
llm, "complete",
|
llm, "complete",
|
||||||
lambda messages, backend=None, model=None:
|
lambda messages, backend=None, model=None, **_:
|
||||||
'{"mood":"focused","valence":0.7,"new_reflections":["I got some thinking done."]}',
|
'{"mood":"focused","valence":0.7,"new_reflections":["I got some thinking done."]}',
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
"""llm.complete: `max_tokens` and `timeout` are threaded into the backend call.
|
||||||
|
|
||||||
|
The OpenAI client is faked so nothing hits a network. We assert the generation
|
||||||
|
cap reaches the create() call and the fast-fail timeout reaches the client (with
|
||||||
|
max_retries=0 so summary.py owns the retry policy, not the SDK).
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import types
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from lyra import llm
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def fake_openai(monkeypatch):
|
||||||
|
recorded: dict = {}
|
||||||
|
|
||||||
|
class FakeCompletions:
|
||||||
|
def create(self, **kwargs):
|
||||||
|
recorded["create"] = kwargs
|
||||||
|
msg = types.SimpleNamespace(content="ok")
|
||||||
|
return types.SimpleNamespace(choices=[types.SimpleNamespace(message=msg)])
|
||||||
|
|
||||||
|
class FakeClient:
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
recorded["client"] = kwargs
|
||||||
|
self.chat = types.SimpleNamespace(completions=FakeCompletions())
|
||||||
|
|
||||||
|
monkeypatch.setattr(llm, "OpenAI", FakeClient)
|
||||||
|
monkeypatch.setattr(llm, "load", lambda: types.SimpleNamespace(
|
||||||
|
mi50_base_url="http://mi50/v1", mi50_model="local-gpu",
|
||||||
|
cloud_model="gpt-4o-mini", openai_api_key="sk-test", local_model="l",
|
||||||
|
))
|
||||||
|
return recorded
|
||||||
|
|
||||||
|
|
||||||
|
def test_mi50_threads_max_tokens_and_timeout(fake_openai):
|
||||||
|
out = llm.complete([{"role": "user", "content": "hi"}],
|
||||||
|
backend="mi50", max_tokens=768, timeout=150)
|
||||||
|
|
||||||
|
assert out == "ok"
|
||||||
|
assert fake_openai["create"]["max_tokens"] == 768
|
||||||
|
assert fake_openai["client"]["timeout"] == 150
|
||||||
|
assert fake_openai["client"]["max_retries"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_cloud_threads_max_tokens_and_timeout(fake_openai):
|
||||||
|
llm.complete([{"role": "user", "content": "hi"}],
|
||||||
|
backend="cloud", max_tokens=768, timeout=150)
|
||||||
|
|
||||||
|
assert fake_openai["create"]["max_tokens"] == 768
|
||||||
|
assert fake_openai["client"]["timeout"] == 150
|
||||||
|
assert fake_openai["client"]["max_retries"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_defaults_omit_cap_and_keep_current_behavior(fake_openai):
|
||||||
|
# No cap / timeout passed -> create() gets no max_tokens, client unbounded.
|
||||||
|
llm.complete([{"role": "user", "content": "hi"}], backend="mi50")
|
||||||
|
|
||||||
|
assert "max_tokens" not in fake_openai["create"]
|
||||||
|
assert "timeout" not in fake_openai["client"]
|
||||||
@@ -0,0 +1,97 @@
|
|||||||
|
"""Summary consolidation: MI50 length cap, fast-fail, and cloud fallback.
|
||||||
|
|
||||||
|
Everything is stubbed — no real backend is touched. These drive the behavior of
|
||||||
|
`summary._summarize_text`: try the primary backend a bounded number of times with
|
||||||
|
a capped generation length, and fall back to cloud if the primary keeps failing.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import types
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from lyra import summary
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def calls(monkeypatch):
|
||||||
|
"""Capture every llm.complete call; per-test behavior via `fake.responder`."""
|
||||||
|
recorded: list[dict] = []
|
||||||
|
|
||||||
|
def fake_complete(messages, backend="local", model=None,
|
||||||
|
max_tokens=None, timeout=None):
|
||||||
|
recorded.append({"backend": backend, "max_tokens": max_tokens, "timeout": timeout})
|
||||||
|
return fake_complete.responder(backend)
|
||||||
|
|
||||||
|
fake_complete.responder = lambda backend: "gist"
|
||||||
|
monkeypatch.setattr(summary.llm, "complete", fake_complete)
|
||||||
|
monkeypatch.setattr(summary.time, "sleep", lambda *_: None) # instant backoff
|
||||||
|
return types.SimpleNamespace(recorded=recorded, fake=fake_complete)
|
||||||
|
|
||||||
|
|
||||||
|
def _set_key(monkeypatch, key="sk-test"):
|
||||||
|
monkeypatch.setattr(summary.config, "load",
|
||||||
|
lambda: types.SimpleNamespace(openai_api_key=key))
|
||||||
|
|
||||||
|
|
||||||
|
def test_falls_back_to_cloud_after_mi50_attempts(calls, monkeypatch):
|
||||||
|
_set_key(monkeypatch)
|
||||||
|
|
||||||
|
def responder(backend):
|
||||||
|
if backend == "mi50":
|
||||||
|
raise RuntimeError("Request timed out.")
|
||||||
|
return "cloud-gist"
|
||||||
|
calls.fake.responder = responder
|
||||||
|
|
||||||
|
out = summary._summarize_text("transcript", "mi50")
|
||||||
|
|
||||||
|
assert out == "cloud-gist"
|
||||||
|
assert [c["backend"] for c in calls.recorded] == ["mi50", "mi50", "cloud"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_fallback_when_backend_is_cloud(calls, monkeypatch):
|
||||||
|
_set_key(monkeypatch)
|
||||||
|
calls.fake.responder = lambda backend: (_ for _ in ()).throw(RuntimeError("boom"))
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
summary._summarize_text("t", "cloud")
|
||||||
|
|
||||||
|
# Cloud is already the primary: retry it, but never a redundant fallback.
|
||||||
|
assert [c["backend"] for c in calls.recorded] == ["cloud", "cloud"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_fallback_without_openai_key(calls, monkeypatch):
|
||||||
|
_set_key(monkeypatch, key="")
|
||||||
|
calls.fake.responder = lambda backend: (_ for _ in ()).throw(RuntimeError("mi50 down"))
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
summary._summarize_text("t", "mi50")
|
||||||
|
|
||||||
|
assert [c["backend"] for c in calls.recorded] == ["mi50", "mi50"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_caps_length_and_timeout_on_every_call(calls, monkeypatch):
|
||||||
|
_set_key(monkeypatch)
|
||||||
|
|
||||||
|
def responder(backend):
|
||||||
|
if backend == "mi50":
|
||||||
|
raise RuntimeError("nope")
|
||||||
|
return "cloud-gist"
|
||||||
|
calls.fake.responder = responder
|
||||||
|
|
||||||
|
summary._summarize_text("t", "mi50")
|
||||||
|
|
||||||
|
assert calls.recorded
|
||||||
|
for c in calls.recorded:
|
||||||
|
assert c["max_tokens"] == summary.SUMMARY_MAX_TOKENS
|
||||||
|
assert c["timeout"] == summary.SUMMARY_TIMEOUT
|
||||||
|
|
||||||
|
|
||||||
|
def test_happy_path_uses_primary_only(calls, monkeypatch):
|
||||||
|
_set_key(monkeypatch)
|
||||||
|
calls.fake.responder = lambda backend: "mi50-gist"
|
||||||
|
|
||||||
|
out = summary._summarize_text("t", "mi50")
|
||||||
|
|
||||||
|
assert out == "mi50-gist"
|
||||||
|
assert [c["backend"] for c in calls.recorded] == ["mi50"] # no retries, no fallback
|
||||||
Reference in New Issue
Block a user