diff --git a/lyra/llm.py b/lyra/llm.py index 080a6cc..f264a8c 100644 --- a/lyra/llm.py +++ b/lyra/llm.py @@ -37,30 +37,45 @@ def _resolved_model(cfg, backend: Backend, model: str | None) -> str: 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 - (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() mdl = _resolved_model(cfg, backend, model) logbus.log("info", "llm call", kind="complete", backend=backend, model=mdl, tok=_approx_tok(messages)) t0 = time.monotonic() - if backend == "cloud": - if not cfg.openai_api_key: - raise RuntimeError("OPENAI_API_KEY is not set") - client = OpenAI(api_key=cfg.openai_api_key) - resp = client.chat.completions.create(model=mdl, messages=messages) - out = resp.choices[0].message.content or "" - elif backend == "mi50": - # MI50 box runs an OpenAI-compatible llama.cpp server; key is unused. - client = OpenAI(api_key="not-needed", base_url=cfg.mi50_base_url) - resp = client.chat.completions.create(model=mdl, messages=messages) + if backend in ("cloud", "mi50"): + if backend == "cloud": + if not cfg.openai_api_key: + raise RuntimeError("OPENAI_API_KEY is not set") + client_kwargs: dict = {"api_key": cfg.openai_api_key} + else: + # MI50 box runs an OpenAI-compatible llama.cpp server; key is unused. + client_kwargs = {"api_key": "not-needed", "base_url": cfg.mi50_base_url} + if timeout is not None: + 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 "" else: + payload: dict = {"model": mdl, "messages": messages, "stream": False} + if max_tokens is not None: + payload["options"] = {"num_predict": max_tokens} resp = httpx.post( f"{cfg.local_base_url}/api/chat", - json={"model": mdl, "messages": messages, "stream": False}, - timeout=120, + json=payload, + timeout=timeout or 120, ) resp.raise_for_status() out = resp.json()["message"]["content"] diff --git a/lyra/summary.py b/lyra/summary.py index 903072b..6057226 100644 --- a/lyra/summary.py +++ b/lyra/summary.py @@ -17,7 +17,16 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from lyra import config, llm, logbus, memory 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. SUMMARIZE_AFTER = 20 @@ -61,16 +70,32 @@ def _summarize_text(text: str, backend: Backend) -> str: {"role": "system", "content": _PROMPT}, {"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: - return llm.complete(messages, backend=backend) + return _call(backend) except Exception as exc: - if attempt == _RETRIES - 1: - raise - logbus.log("debug", "summary retry", attempt=attempt + 1, error=str(exc)[:80]) - time.sleep(5 * (attempt + 1)) - raise RuntimeError("unreachable") + last_exc = exc + logbus.log("debug", "summary retry", attempt=attempt + 1, + backend=backend, error=str(exc)[:80]) + if attempt < MI50_ATTEMPTS - 1: + 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: diff --git a/tests/test_dream.py b/tests/test_dream.py index 867db3d..81a2c14 100644 --- a/tests/test_dream.py +++ b/tests/test_dream.py @@ -20,7 +20,7 @@ def lyra(tmp_path, monkeypatch): # reflect() expects JSON back; everything else just stores the text. monkeypatch.setattr( 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."]}', ) diff --git a/tests/test_llm_bounds.py b/tests/test_llm_bounds.py new file mode 100644 index 0000000..a7e26d4 --- /dev/null +++ b/tests/test_llm_bounds.py @@ -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"] diff --git a/tests/test_summary_fallback.py b/tests/test_summary_fallback.py new file mode 100644 index 0000000..df4ee0d --- /dev/null +++ b/tests/test_summary_fallback.py @@ -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