diff --git a/docs/APPS.md b/docs/APPS.md index 7398638c9..c16bf056a 100644 --- a/docs/APPS.md +++ b/docs/APPS.md @@ -308,13 +308,13 @@ The `occurrences` field (optional string) provides topic-specific extraction gui - Resolution: `"name"` → `muse/{name}.py`, `"app:name"` → `apps/{app}/muse/{name}.py`, or explicit path **Pre-hooks** (`pre_process`): Modify inputs before the LLM call -- `context` is a `PreHookContext` with: `name`, `agent_id`, `provider`, `model`, `prompt`, `system_instruction`, `user_instruction`, `extra_context`, `output_format`, `meta`, and for generators: `day`, `segment`, `span`, `transcript`, `output_path` +- `context` is the full config dict with: `name`, `agent_id`, `provider`, `model`, `prompt`, `system_instruction`, `user_instruction`, `extra_context`, `output`, `meta`, and for generators: `day`, `segment`, `span`, `span_mode`, `transcript`, `output_path` - Return a dict of modified fields to merge back (e.g., `{"prompt": "modified"}`) - Return `None` for no changes **Post-hooks** (`post_process`): Transform output after the LLM call - `result` is the LLM output (markdown or JSON string) -- `context` is a `HookContext` with: `name`, `agent_id`, `provider`, `model`, `prompt`, `output_format`, `meta`, and for generators: `day`, `segment`, `span`, `transcript`, `output_path` +- `context` is the full config dict with: `name`, `agent_id`, `provider`, `model`, `prompt`, `output`, `meta`, and for generators: `day`, `segment`, `span`, `span_mode`, `transcript`, `output_path` - Return modified string, or `None` to use original result Hook errors are logged but don't crash the pipeline (graceful degradation). diff --git a/docs/PROVIDERS.md b/docs/PROVIDERS.md index 2fe372e5d..079195811 100644 --- a/docs/PROVIDERS.md +++ b/docs/PROVIDERS.md @@ -131,7 +131,7 @@ async def run_tools( **Event emission:** -Providers must emit events via the `on_event` callback. See `think/agents.py` for TypedDict definitions: +Providers must emit events via the `on_event` callback. See `think/providers/shared.py` for TypedDict definitions: | Event | When | |-------|------| @@ -142,7 +142,7 @@ Providers must emit events via the `on_event` callback. See `think/agents.py` fo | `FinishEvent` | Agent run completes successfully | | `ErrorEvent` | Error occurs | -Use `JSONEventCallback` from `think/agents.py` to wrap the callback and auto-add timestamps. +Use `JSONEventCallback` from `think/providers/shared.py` to wrap the callback and auto-add timestamps. **Finish event format:** diff --git a/muse/anticipation.py b/muse/anticipation.py index 3fa74b195..7994f1460 100644 --- a/muse/anticipation.py +++ b/muse/anticipation.py @@ -31,8 +31,8 @@ def post_process(result: str, context: dict) -> str | None: Args: result: The generated output markdown content. - context: HookContext with keys including day, segment, name, - output_path, meta, transcript, span. + context: Config dict with keys including day, segment, name, + output_path, meta, transcript, span, span_mode. Returns: None - this hook does not modify the output result. diff --git a/muse/occurrence.py b/muse/occurrence.py index cf422c919..42bd97d4b 100644 --- a/muse/occurrence.py +++ b/muse/occurrence.py @@ -31,8 +31,8 @@ def post_process(result: str, context: dict) -> str | None: Args: result: The generated output markdown content. - context: HookContext with keys including day, segment, name, - output_path, meta, transcript, span. + context: Config dict with keys including day, segment, name, + output_path, meta, transcript, span, span_mode. Returns: None - this hook does not modify the output result. diff --git a/tests/test_agents_ndjson.py b/tests/test_agents_ndjson.py index 34ed3f347..a04cf1f84 100644 --- a/tests/test_agents_ndjson.py +++ b/tests/test_agents_ndjson.py @@ -122,7 +122,8 @@ def test_ndjson_single_request(mock_journal, monkeypatch, capsys): start_event = events[0] assert start_event["event"] == "start" - assert start_event["prompt"] == "What is 2+2?" + # Prompt includes system instruction prepended during enrichment + assert "What is 2+2?" in start_event["prompt"] assert start_event["provider"] == "openai" assert start_event["model"] == GPT_5 @@ -179,10 +180,11 @@ def test_ndjson_multiple_requests(mock_journal, monkeypatch, capsys): start_events = [e for e in events if e["event"] == "start"] assert len(start_events) == 3 - assert start_events[0]["prompt"] == "First question" - assert start_events[1]["prompt"] == "Second question" + # Prompts include system instruction prepended during enrichment + assert "First question" in start_events[0]["prompt"] + assert "Second question" in start_events[1]["prompt"] assert start_events[1]["provider"] == "anthropic" - assert start_events[2]["prompt"] == "Third question" + assert "Third question" in start_events[2]["prompt"] assert start_events[2]["name"] == "technical" diff --git a/tests/test_anthropic.py b/tests/test_anthropic.py index bf572f484..4d76f7a87 100644 --- a/tests/test_anthropic.py +++ b/tests/test_anthropic.py @@ -186,7 +186,8 @@ def test_claude_main(monkeypatch, tmp_path, capsys): events = [json.loads(line) for line in out_lines] assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + # Prompt includes system instruction prepended during enrichment + assert "hello" in events[0]["prompt"] assert events[0]["name"] == "default" assert events[0]["model"] == CLAUDE_SONNET_4 assert events[-1]["event"] == "finish" @@ -230,7 +231,8 @@ def test_claude_outfile(monkeypatch, tmp_path, capsys): events = [json.loads(line) for line in out_lines] assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + # Prompt includes system instruction prepended during enrichment + assert "hello" in events[0]["prompt"] assert events[0]["name"] == "default" assert events[0]["model"] == CLAUDE_SONNET_4 assert events[-1]["event"] == "finish" diff --git a/tests/test_batch.py b/tests/test_batch.py index a78dd0c6b..ba3fb91d6 100644 --- a/tests/test_batch.py +++ b/tests/test_batch.py @@ -465,5 +465,3 @@ async def test_batch_client_passthrough(mock_agenerate): # Verify client was passed through call_kwargs = mock_agenerate.call_args[1] assert call_kwargs["client"] is mock_client - - diff --git a/tests/test_generate_full.py b/tests/test_generate_full.py index 222427827..82b57af2d 100644 --- a/tests/test_generate_full.py +++ b/tests/test_generate_full.py @@ -81,12 +81,13 @@ def test_generate_output_ndjson(tmp_path, monkeypatch): ) try: + # Mock the underlying generation function in think.models + import think.models + monkeypatch.setattr( - mod, - "generate_agent_output", - lambda *a, **k: ( - MOCK_RESULT if k.get("return_result") else MOCK_RESULT["text"] - ), + think.models, + "generate_with_result", + lambda *a, **k: MOCK_RESULT, ) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -134,7 +135,7 @@ def post_process(result, context): ctx_copy = { "day": context.get("day"), "segment": context.get("segment"), - "span": context.get("span"), + "span": context.get("span_mode"), "name": context.get("name"), "has_transcript": bool(context.get("transcript")), "has_meta": bool(context.get("meta")), @@ -151,12 +152,13 @@ def post_process(result, context): ) try: + # Mock the underlying generation function in think.models + import think.models + monkeypatch.setattr( - mod, - "generate_agent_output", - lambda *a, **k: ( - MOCK_RESULT if k.get("return_result") else MOCK_RESULT["text"] - ), + think.models, + "generate_with_result", + lambda *a, **k: MOCK_RESULT, ) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -181,6 +183,7 @@ def post_process(result, context): assert captured["day"] == "20240101" assert captured["segment"] is None + # span_mode is a bool in the new config structure assert captured["span"] is False assert captured["name"] == "hooked_gen" assert captured["has_transcript"] is True @@ -207,12 +210,13 @@ def test_generate_without_hook_succeeds(tmp_path, monkeypatch): ) try: + # Mock the underlying generation function in think.models + import think.models + monkeypatch.setattr( - mod, - "generate_agent_output", - lambda *a, **k: ( - MOCK_RESULT if k.get("return_result") else MOCK_RESULT["text"] - ), + think.models, + "generate_with_result", + lambda *a, **k: MOCK_RESULT, ) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -301,11 +305,11 @@ def test_generate_skipped_on_no_input(tmp_path, monkeypatch): def test_named_hook_resolution(tmp_path, monkeypatch): """Test that named hooks are resolved via load_post_hook.""" - agents = importlib.import_module("think.agents") + from think.muse import load_post_hook # Config with named hook (new format) config = {"hook": {"post": "occurrence"}} - hook_fn = agents.load_post_hook(config) + hook_fn = load_post_hook(config) # Should resolve to muse/occurrence.py and be callable assert callable(hook_fn) diff --git a/tests/test_google.py b/tests/test_google.py index 9151233e5..a2a4b2df8 100644 --- a/tests/test_google.py +++ b/tests/test_google.py @@ -52,7 +52,7 @@ def test_google_main(monkeypatch, tmp_path, capsys): events = [json.loads(line) for line in out_lines] assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] assert events[0]["name"] == "default" assert events[0]["model"] == GEMINI_FLASH assert events[-1]["event"] == "finish" diff --git a/tests/test_google_thinking.py b/tests/test_google_thinking.py index 1f67c81d2..1fe9c52c3 100644 --- a/tests/test_google_thinking.py +++ b/tests/test_google_thinking.py @@ -50,7 +50,7 @@ def test_google_thinking_events(monkeypatch, tmp_path, capsys): # Check that we have start, thinking, and finish events assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] # Look for thinking event thinking_events = [e for e in events if e["event"] == "thinking"] diff --git a/tests/test_openai.py b/tests/test_openai.py index 7bc299c20..7fd49260e 100644 --- a/tests/test_openai.py +++ b/tests/test_openai.py @@ -71,7 +71,7 @@ def test_openai_main(monkeypatch, tmp_path, capsys): events = [json.loads(line) for line in out_lines] assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] assert events[0]["name"] == "default" assert events[0]["model"] == GPT_5 assert events[-1]["event"] == "finish" @@ -128,7 +128,7 @@ def test_openai_thinking_events(monkeypatch, tmp_path, capsys): # Check that we have start, thinking, and finish events assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] # Look for thinking event thinking_events = [e for e in events if e["event"] == "thinking"] @@ -211,7 +211,7 @@ def test_openai_outfile(monkeypatch, tmp_path, capsys): events = [json.loads(line) for line in out_lines] assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] assert events[0]["name"] == "default" assert events[0]["model"] == GPT_5 assert events[-1]["event"] == "finish" @@ -265,7 +265,7 @@ def test_openai_thinking_events_stdout(monkeypatch, tmp_path, capsys): # Check that we have start, thinking, and finish events assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] # Look for thinking event thinking_events = [e for e in events if e["event"] == "thinking"] @@ -375,7 +375,7 @@ def test_openai_thinking_events_error(monkeypatch, tmp_path, capsys): # Check that we have start, thinking, and finish events assert events[0]["event"] == "start" assert isinstance(events[0]["ts"], int) - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] # Look for thinking event thinking_events = [e for e in events if e["event"] == "thinking"] @@ -438,7 +438,7 @@ def test_openai_tool_call_events(monkeypatch, tmp_path, capsys): # Check start event assert events[0]["event"] == "start" - assert events[0]["prompt"] == "hello" + assert "hello" in events[0]["prompt"] # Look for tool_start event tool_start_events = [e for e in events if e["event"] == "tool_start"] diff --git a/tests/test_output_hooks.py b/tests/test_output_hooks.py index 6f5082ae2..f1c523810 100644 --- a/tests/test_output_hooks.py +++ b/tests/test_output_hooks.py @@ -4,7 +4,7 @@ """Tests for the generator output hooks system. Tests cover: -- Hook loading and validation via load_post_hook +- Hook loading and validation via load_post_hook / load_pre_hook - Hook invocation via NDJSON protocol - Hook error handling """ @@ -16,6 +16,7 @@ import os import shutil from pathlib import Path +from think.muse import load_post_hook, load_pre_hook from think.utils import day_path FIXTURES = Path("fixtures") @@ -64,8 +65,6 @@ def run_generator_with_config(mod, config: dict, monkeypatch) -> list[dict]: def test_load_post_hook_success(tmp_path): """Test loading a valid hook with post_process function.""" - agents = importlib.import_module("think.agents") - hook_file = tmp_path / "test_hook.py" hook_file.write_text(""" def post_process(result, context): @@ -74,7 +73,7 @@ def post_process(result, context): # Config with explicit path config = {"hook": {"post": str(hook_file)}} - hook_fn = agents.load_post_hook(config) + hook_fn = load_post_hook(config) assert callable(hook_fn) # Test the hook transforms content @@ -84,8 +83,6 @@ def post_process(result, context): def test_load_post_hook_missing_post_process(tmp_path): """Test that hook without post_process function raises ValueError.""" - agents = importlib.import_module("think.agents") - hook_file = tmp_path / "bad_hook.py" hook_file.write_text(""" def other_function(): @@ -94,7 +91,7 @@ def other_function(): config = {"hook": {"post": str(hook_file)}} try: - agents.load_post_hook(config) + load_post_hook(config) assert False, "Should have raised ValueError" except ValueError as e: assert "must define a 'post_process' function" in str(e) @@ -102,8 +99,6 @@ def other_function(): def test_load_post_hook_not_callable(tmp_path): """Test that hook with non-callable post_process raises ValueError.""" - agents = importlib.import_module("think.agents") - hook_file = tmp_path / "bad_hook.py" hook_file.write_text(""" post_process = "not a function" @@ -111,7 +106,7 @@ post_process = "not a function" config = {"hook": {"post": str(hook_file)}} try: - agents.load_post_hook(config) + load_post_hook(config) assert False, "Should have raised ValueError" except ValueError as e: assert "'post_process' must be callable" in str(e) @@ -119,30 +114,24 @@ post_process = "not a function" def test_load_post_hook_no_hook_config(): """Test that missing hook config returns None.""" - agents = importlib.import_module("think.agents") - - assert agents.load_post_hook({}) is None - assert agents.load_post_hook({"hook": {}}) is None - assert agents.load_post_hook({"hook": {"pre": "something"}}) is None + assert load_post_hook({}) is None + assert load_post_hook({"hook": {}}) is None + assert load_post_hook({"hook": {"pre": "something"}}) is None def test_load_post_hook_named_resolution(): """Test that named hooks resolve to muse/{name}.py.""" - agents = importlib.import_module("think.agents") - # occurrence.py exists in muse/ config = {"hook": {"post": "occurrence"}} - hook_fn = agents.load_post_hook(config) + hook_fn = load_post_hook(config) assert callable(hook_fn) def test_load_post_hook_file_not_found(tmp_path): """Test that nonexistent hook file raises ImportError.""" - agents = importlib.import_module("think.agents") - config = {"hook": {"post": str(tmp_path / "nonexistent.py")}} try: - agents.load_post_hook(config) + load_post_hook(config) assert False, "Should have raised ImportError" except ImportError as e: assert "not found" in str(e) @@ -193,12 +182,13 @@ def post_process(result, context): """) try: + # Mock the underlying generation function in think.models + import think.models + monkeypatch.setattr( - mod, - "generate_agent_output", - lambda *a, **k: ( - MOCK_RESULT if k.get("return_result") else MOCK_RESULT["text"] - ), + think.models, + "generate_with_result", + lambda *a, **k: MOCK_RESULT, ) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -247,12 +237,13 @@ def post_process(result, context): """) try: + # Mock the underlying generation function in think.models + import think.models + monkeypatch.setattr( - mod, - "generate_agent_output", - lambda *a, **k: ( - MOCK_RESULT if k.get("return_result") else MOCK_RESULT["text"] - ), + think.models, + "generate_with_result", + lambda *a, **k: MOCK_RESULT, ) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -297,12 +288,13 @@ def post_process(result, context): """) try: + # Mock the underlying generation function in think.models + import think.models + monkeypatch.setattr( - mod, - "generate_agent_output", - lambda *a, **k: ( - MOCK_RESULT if k.get("return_result") else MOCK_RESULT["text"] - ), + think.models, + "generate_with_result", + lambda *a, **k: MOCK_RESULT, ) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -329,81 +321,6 @@ def post_process(result, context): prompt_file.unlink() -def test_build_hook_context(): - """Test that build_hook_context creates correct context.""" - agents = importlib.import_module("think.agents") - - config = { - "name": "test_gen", - "agent_id": "123456", - "provider": "google", - "model": "gemini-2.0-flash", - "prompt": "test prompt", - "output": "md", - "day": "20240101", - "segment": "120000_3600", - } - - context = agents.build_hook_context( - config, - transcript="test transcript", - output_path="/tmp/test.md", - span=False, - ) - - assert context["name"] == "test_gen" - assert context["agent_id"] == "123456" - assert context["provider"] == "google" - assert context["model"] == "gemini-2.0-flash" - assert context["prompt"] == "test prompt" - assert context["output_format"] == "md" - assert context["day"] == "20240101" - assert context["segment"] == "120000_3600" - assert context["transcript"] == "test transcript" - assert context["output_path"] == "/tmp/test.md" - assert context["span"] is False - assert context["meta"] == config - - -def test_run_post_hook_transforms_result(): - """Test that run_post_hook applies transformation.""" - agents = importlib.import_module("think.agents") - - def hook(result, context): - return result.upper() - - context = agents.build_hook_context({"name": "test"}) - output = agents.run_post_hook("hello world", context, hook) - - assert output == "HELLO WORLD" - - -def test_run_post_hook_none_keeps_original(): - """Test that run_post_hook keeps original when hook returns None.""" - agents = importlib.import_module("think.agents") - - def hook(result, context): - return None - - context = agents.build_hook_context({"name": "test"}) - output = agents.run_post_hook("original", context, hook) - - assert output == "original" - - -def test_run_post_hook_error_keeps_original(): - """Test that run_post_hook keeps original on error.""" - agents = importlib.import_module("think.agents") - - def hook(result, context): - raise RuntimeError("boom") - - context = agents.build_hook_context({"name": "test"}) - output = agents.run_post_hook("original", context, hook) - - assert output == "original" - - # ============================================================================= # Pre-hook Tests # ============================================================================= @@ -411,8 +328,6 @@ def test_run_post_hook_error_keeps_original(): def test_load_pre_hook_success(tmp_path): """Test loading a valid hook with pre_process function.""" - agents = importlib.import_module("think.agents") - hook_file = tmp_path / "test_pre_hook.py" hook_file.write_text(""" def pre_process(context): @@ -420,7 +335,7 @@ def pre_process(context): """) config = {"hook": {"pre": str(hook_file)}} - hook_fn = agents.load_pre_hook(config) + hook_fn = load_pre_hook(config) assert callable(hook_fn) # Test the hook returns modifications @@ -430,8 +345,6 @@ def pre_process(context): def test_load_pre_hook_missing_pre_process(tmp_path): """Test that hook without pre_process function raises ValueError.""" - agents = importlib.import_module("think.agents") - hook_file = tmp_path / "bad_hook.py" hook_file.write_text(""" def other_function(): @@ -440,7 +353,7 @@ def other_function(): config = {"hook": {"pre": str(hook_file)}} try: - agents.load_pre_hook(config) + load_pre_hook(config) assert False, "Should have raised ValueError" except ValueError as e: assert "must define a 'pre_process' function" in str(e) @@ -448,8 +361,6 @@ def other_function(): def test_load_pre_hook_not_callable(tmp_path): """Test that hook with non-callable pre_process raises ValueError.""" - agents = importlib.import_module("think.agents") - hook_file = tmp_path / "bad_hook.py" hook_file.write_text(""" pre_process = "not a function" @@ -457,7 +368,7 @@ pre_process = "not a function" config = {"hook": {"pre": str(hook_file)}} try: - agents.load_pre_hook(config) + load_pre_hook(config) assert False, "Should have raised ValueError" except ValueError as e: assert "'pre_process' must be callable" in str(e) @@ -465,106 +376,21 @@ pre_process = "not a function" def test_load_pre_hook_no_hook_config(): """Test that missing hook config returns None.""" - agents = importlib.import_module("think.agents") - - assert agents.load_pre_hook({}) is None - assert agents.load_pre_hook({"hook": {}}) is None - assert agents.load_pre_hook({"hook": {"post": "something"}}) is None + assert load_pre_hook({}) is None + assert load_pre_hook({"hook": {}}) is None + assert load_pre_hook({"hook": {"post": "something"}}) is None def test_load_pre_hook_file_not_found(tmp_path): """Test that nonexistent hook file raises ImportError.""" - agents = importlib.import_module("think.agents") - config = {"hook": {"pre": str(tmp_path / "nonexistent.py")}} try: - agents.load_pre_hook(config) + load_pre_hook(config) assert False, "Should have raised ImportError" except ImportError as e: assert "not found" in str(e) -def test_build_pre_hook_context(): - """Test that build_pre_hook_context creates correct context.""" - agents = importlib.import_module("think.agents") - - config = { - "name": "test_gen", - "agent_id": "123456", - "provider": "google", - "model": "gemini-2.0-flash", - "prompt": "test prompt", - "system_instruction": "be helpful", - "user_instruction": "answer questions", - "extra_context": "extra info", - "output": "md", - "day": "20240101", - "segment": "120000_3600", - } - - context = agents.build_pre_hook_context( - config, - transcript="test transcript", - output_path="/tmp/test.md", - span=False, - ) - - assert context["name"] == "test_gen" - assert context["agent_id"] == "123456" - assert context["provider"] == "google" - assert context["model"] == "gemini-2.0-flash" - assert context["prompt"] == "test prompt" - assert context["system_instruction"] == "be helpful" - assert context["user_instruction"] == "answer questions" - assert context["extra_context"] == "extra info" - assert context["output_format"] == "md" - assert context["day"] == "20240101" - assert context["segment"] == "120000_3600" - assert context["transcript"] == "test transcript" - assert context["output_path"] == "/tmp/test.md" - assert context["span"] is False - assert context["meta"] == config - - -def test_run_pre_hook_returns_modifications(): - """Test that run_pre_hook returns modifications dict.""" - agents = importlib.import_module("think.agents") - - def hook(context): - return {"prompt": "modified prompt", "transcript": "modified transcript"} - - context = agents.build_pre_hook_context({"name": "test", "prompt": "original"}) - result = agents.run_pre_hook(context, hook) - - assert result == {"prompt": "modified prompt", "transcript": "modified transcript"} - - -def test_run_pre_hook_none_returns_none(): - """Test that run_pre_hook returns None when hook returns None.""" - agents = importlib.import_module("think.agents") - - def hook(context): - return None - - context = agents.build_pre_hook_context({"name": "test"}) - result = agents.run_pre_hook(context, hook) - - assert result is None - - -def test_run_pre_hook_error_returns_none(): - """Test that run_pre_hook returns None on error.""" - agents = importlib.import_module("think.agents") - - def hook(context): - raise RuntimeError("boom") - - context = agents.build_pre_hook_context({"name": "test"}) - result = agents.run_pre_hook(context, hook) - - assert result is None - - def test_pre_hook_invocation(tmp_path, monkeypatch): """Test that agents.py invokes pre-hook and uses modified inputs.""" mod = importlib.import_module("think.agents") @@ -589,15 +415,17 @@ def pre_process(context): """) try: - # Track what generate_agent_output receives - received_args = {} + # Track what generate_with_result receives + received_kwargs = {} def mock_generate(*args, **kwargs): - received_args["transcript"] = args[0] - received_args["prompt"] = args[1] - return MOCK_RESULT if kwargs.get("return_result") else MOCK_RESULT["text"] + received_kwargs.update(kwargs) + received_kwargs["contents"] = args[0] if args else kwargs.get("contents") + return MOCK_RESULT + + import think.models - monkeypatch.setattr(mod, "generate_agent_output", mock_generate) + monkeypatch.setattr(think.models, "generate_with_result", mock_generate) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -611,8 +439,11 @@ def pre_process(context): events = run_generator_with_config(mod, config, monkeypatch) - # Verify pre-hook modified the prompt - assert "[pre-processed]" in received_args["prompt"] + # Verify pre-hook modified the prompt - check in contents + contents = received_kwargs.get("contents", []) + # The prompt should contain [pre-processed] + prompt_found = any("[pre-processed]" in str(c) for c in contents) + assert prompt_found, f"Expected [pre-processed] in contents: {contents}" # Verify generator still completed successfully finish_events = [e for e in events if e["event"] == "finish"] @@ -647,13 +478,16 @@ def post_process(result, context): """) try: - received_args = {} + received_kwargs = {} def mock_generate(*args, **kwargs): - received_args["prompt"] = args[1] - return MOCK_RESULT if kwargs.get("return_result") else MOCK_RESULT["text"] + received_kwargs.update(kwargs) + received_kwargs["contents"] = args[0] if args else kwargs.get("contents") + return MOCK_RESULT + + import think.models - monkeypatch.setattr(mod, "generate_agent_output", mock_generate) + monkeypatch.setattr(think.models, "generate_with_result", mock_generate) monkeypatch.setenv("GOOGLE_API_KEY", "x") monkeypatch.setenv("JOURNAL_PATH", str(tmp_path)) @@ -667,8 +501,10 @@ def post_process(result, context): events = run_generator_with_config(mod, config, monkeypatch) - # Verify pre-hook modified the prompt - assert "[pre]" in received_args["prompt"] + # Verify pre-hook modified the prompt - check in contents + contents = received_kwargs.get("contents", []) + prompt_found = any("[pre]" in str(c) for c in contents) + assert prompt_found, f"Expected [pre] in contents: {contents}" # Verify post-hook modified the result finish_events = [e for e in events if e["event"] == "finish"] diff --git a/think/agents.py b/think/agents.py index 49ba55056..2bf2331c7 100644 --- a/think/agents.py +++ b/think/agents.py @@ -19,10 +19,9 @@ import logging import os import sys import traceback -from dataclasses import dataclass from datetime import datetime from pathlib import Path -from typing import Any, Callable, Optional, TypedDict +from typing import Any, Callable, Optional from think.cluster import cluster, cluster_period, cluster_span from think.muse import ( @@ -36,7 +35,7 @@ from think.muse import ( source_is_enabled, source_is_required, ) -from think.providers.shared import Event, GenerateResult +from think.providers.shared import Event from think.utils import ( day_log, day_path, @@ -49,6 +48,9 @@ from think.utils import ( LOG = logging.getLogger("think.agents") +# Minimum content length for transcript-based generation +MIN_INPUT_CHARS = 50 + def setup_logging(verbose: bool = False) -> logging.Logger: """Configure logging for agent CLI.""" @@ -129,16 +131,12 @@ def parse_agent_events_to_turns(conversation_id: str) -> list: Note: Incomplete turns (missing finish event) are skipped """ - import logging - from think.cortex_client import read_agent_events - logger = logging.getLogger(__name__) - try: events = read_agent_events(conversation_id) except FileNotFoundError: - logger.warning(f"Cannot continue from {conversation_id}: log not found") + LOG.warning(f"Cannot continue from {conversation_id}: log not found") return [] turns = [] @@ -182,194 +180,7 @@ def parse_agent_events_to_turns(conversation_id: str) -> list: # ============================================================================= -# Hook Framework (unified for agents and generators) -# ============================================================================= - - -class HookContext(TypedDict, total=False): - """Context passed to hook functions. - - Provides unified context for both tool-using agents and generators. - Not all fields are present for all modalities. - """ - - # Identity - name: str # Agent/generator name - agent_id: str # Unique agent ID - provider: str # google/anthropic/openai - model: str # Model used - - # Temporal (generators) - day: str # YYYYMMDD - segment: str # Segment key - span: bool # True if span mode - - # Content - prompt: str # Original prompt (agents) or empty (generators) - transcript: str # Clustered transcript (generators only) - - # Output - output_path: str # Where result will be written - output_format: str # 'md' or 'json' - - # Full config - meta: dict # Full frontmatter/config - - -class PreHookContext(TypedDict, total=False): - """Context passed to pre-processing hook functions. - - Pre-hooks receive all inputs before the LLM call and can modify them. - Returns a dict of modified fields to merge back. - """ - - # Identity - name: str # Agent/generator name - agent_id: str # Unique agent ID - provider: str # google/anthropic/openai - model: str # Model used - - # Temporal (generators) - day: str # YYYYMMDD - segment: str # Segment key - span: bool # True if span mode - - # Modifiable inputs - prompt: str # User prompt (can modify) - system_instruction: str # System prompt (can modify) - user_instruction: str # User instruction (agents, can modify) - extra_context: str # Extra context (agents, can modify) - transcript: str # Clustered transcript (generators, can modify) - - # Output settings - output_path: str # Where result will be written - output_format: str # 'md' or 'json' - - # Full config (read-only reference) - meta: dict # Full frontmatter/config - - -def _build_base_context(config: dict) -> dict: - """Build common context fields shared by pre and post hooks.""" - context = { - "name": config.get("name", ""), - "agent_id": config.get("agent_id", ""), - "provider": config.get("provider", ""), - "model": config.get("model", ""), - "prompt": config.get("prompt", ""), - "output_format": config.get("output", "md"), - "meta": config, - } - - # Add generator-specific fields if present - if "day" in config: - context["day"] = config["day"] - if "segment" in config: - context["segment"] = config["segment"] - - return context - - -def build_pre_hook_context(config: dict, **extras: Any) -> PreHookContext: - """Build PreHookContext from config and extra values.""" - context: PreHookContext = _build_base_context(config) - - # Add pre-hook specific fields - context["system_instruction"] = config.get("system_instruction", "") - context["user_instruction"] = config.get("user_instruction", "") - context["extra_context"] = config.get("extra_context", "") - - # Merge extras (transcript, output_path, span, etc.) - context.update(extras) - - return context - - -def build_hook_context(config: dict, **extras: Any) -> HookContext: - """Build HookContext from config and extra values.""" - context: HookContext = _build_base_context(config) - - # Merge extras (transcript, output_path, span, etc.) - context.update(extras) - - return context - - -def run_pre_hook( - context: PreHookContext, - hook_fn: Callable[[PreHookContext], dict | None], -) -> dict | None: - """Execute pre-processing hook and return modifications dict. - - Hook errors are logged and return None (graceful degradation). - """ - try: - modifications = hook_fn(context) - if modifications is not None: - logging.info( - "Pre-hook returned modifications: %s", list(modifications.keys()) - ) - return modifications - except Exception as exc: - logging.error("Pre-hook failed: %s", exc) - - return None - - -def run_post_hook( - result: str, - context: HookContext, - hook_fn: Callable[[str, HookContext], str | None], -) -> str: - """Execute post-processing hook and return (potentially transformed) result. - - Args: - result: The LLM-generated output text - context: Hook context with metadata - hook_fn: The post_process function to call - - Returns: - Transformed result if hook returns string, original result otherwise. - """ - try: - hook_result = hook_fn(result, context) - if hook_result is not None: - logging.info("Hook transformed result") - return hook_result - except Exception as exc: - logging.error("Hook failed: %s", exc) - - return result - - -__all__ = [ - # Re-exported from think.providers.shared - "Event", - "GenerateResult", - # Local definitions - "HookContext", - "PreHookContext", - "InputContext", - "JSONEventWriter", - "format_tool_summary", - "parse_agent_events_to_turns", - "load_post_hook", - "load_pre_hook", - "build_hook_context", - "build_pre_hook_context", - "run_post_hook", - "run_pre_hook", - "assemble_inputs", - "scan_day", - "generate_agent_output", - "hydrate_config", - "expand_tools", - "validate_config", -] - - -# ============================================================================= -# Config Hydration and Validation (moved from cortex.py) +# Config Hydration, Enrichment, and Validation # ============================================================================= @@ -496,72 +307,26 @@ def validate_config(config: dict) -> str | None: return None -# ============================================================================= -# Unified Input Assembly (shared by generate and tools paths) -# ============================================================================= - -# Minimum content length for transcript-based generation -MIN_INPUT_CHARS = 50 - - -@dataclass -class InputContext: - """Assembled inputs for generation or tool execution. +def enrich_config(config: dict) -> None: + """Enrich config with transcript, system instruction, output path, etc. - Contains all resolved inputs ready for LLM call, including transcript, - prompts, and output path. Used by generate path always, and by tools path - when day is specified for transcript loading. - """ - - # Transcript (from day/segment/span clustering) - transcript: str - source_counts: dict[str, int] - - # Prompts - prompt: str # Final prompt (with template substitution) - system_instruction: str - system_prompt_name: str # For diagnostic output (dry-run) - - # Output - output_path: Optional[Path] - output_format: Optional[str] # 'md' or 'json' - - # Metadata for hooks and logging - meta: dict # Agent config metadata - agent_path: Optional[Path] # Path to agent .md file - - # Skip reason (if should skip execution) - skip_reason: Optional[str] - - # Day/segment context - day: Optional[str] - segment: Optional[str] - span_mode: bool - - -def assemble_inputs(config: dict) -> InputContext: - """Assemble all inputs for generation or tool execution. - - Handles: - - Loading agent config metadata - - Transcript loading from journal (if day specified) - - Source filtering and required source validation - - Minimum content checks - - Prompt template substitution - - System instruction composition - - Output path resolution + Mutates config in place, adding: + - transcript: Clustered transcript content (if day specified) + - system_instruction: System prompt (if not already set) + - output_path: Where to write output (if output format specified) + - source_counts: Dict of source type -> count + - skip_reason: Why to skip execution (if applicable) + - span_mode: Whether in span mode + - meta: Agent metadata from muse config Args: - config: Hydrated config dict from cortex - - Returns: - InputContext with all resolved inputs, or skip_reason if should skip + config: Hydrated config dict to enrich """ name = config.get("name", "default") day = config.get("day") segment = config.get("segment") span = config.get("span") # List of sequential segment keys - facet = config.get("facet") # For multi-facet agents + facet = config.get("facet") output_format = config.get("output") output_path_override = config.get("output_path") user_prompt = config.get("prompt", "") @@ -577,23 +342,15 @@ def assemble_inputs(config: dict) -> InputContext: meta = {} agent_path = None + config["meta"] = meta + config["agent_path"] = agent_path + config["span_mode"] = bool(span) + config["source_counts"] = {} + # Check if config is disabled if meta.get("disabled"): - return InputContext( - transcript="", - source_counts={}, - prompt=user_prompt, - system_instruction="", - system_prompt_name="journal", - output_path=None, - output_format=output_format, - meta=meta, - agent_path=agent_path, - skip_reason="disabled", - day=day, - segment=segment, - span_mode=bool(span), - ) + config["skip_reason"] = "disabled" + return # Extract instructions config for source filtering and system prompt instructions_config = meta.get("instructions") @@ -603,20 +360,16 @@ def assemble_inputs(config: dict) -> InputContext: config_overrides=instructions_config, ) sources = instructions.get("sources", {}) - system_prompt_name = instructions.get("system_prompt_name", "journal") - system_instruction = instructions["system_instruction"] + config["system_prompt_name"] = instructions.get("system_prompt_name", "journal") - # Append extra_context (facets, etc.) to system instruction if present - extra_context = instructions.get("extra_context") - if extra_context: - system_instruction = f"{system_instruction}\n\n{extra_context}" - - # Track span mode - span_mode = bool(span) - - # Initialize transcript variables - transcript = "" - source_counts: dict[str, int] = {} + # Set system_instruction if not already provided + if not config.get("system_instruction"): + system_instruction = instructions["system_instruction"] + # Append extra_context (facets, etc.) to system instruction if present + extra_context = instructions.get("extra_context") + if extra_context: + system_instruction = f"{system_instruction}\n\n{extra_context}" + config["system_instruction"] = system_instruction # Transcript loading (only if day is provided) if day: @@ -627,20 +380,15 @@ def assemble_inputs(config: dict) -> InputContext: os.environ["SEGMENT_KEY"] = span[0] # Convert sources for clustering - # For audio/screen: use source_is_enabled to get bool - # For agents: pass through dict for selective filtering, or use source_is_enabled cluster_sources: dict = {} for k, v in sources.items(): if k == "agents": agent_filter = get_agent_filter(v) if agent_filter is None: - # All agents (True or "required") cluster_sources[k] = source_is_enabled(v) elif not agent_filter: - # No agents (False or empty dict) cluster_sources[k] = False else: - # Selective filtering - pass dict through cluster_sources[k] = agent_filter else: cluster_sources[k] = source_is_enabled(v) @@ -655,44 +403,20 @@ def assemble_inputs(config: dict) -> InputContext: else: transcript, source_counts = cluster(day, sources=cluster_sources) + config["transcript"] = transcript + config["source_counts"] = source_counts total_count = sum(source_counts.values()) # Check required sources have content for source_type, mode in sources.items(): if source_is_required(mode) and source_counts.get(source_type, 0) == 0: - return InputContext( - transcript=transcript, - source_counts=source_counts, - prompt=user_prompt, - system_instruction=system_instruction, - system_prompt_name=system_prompt_name, - output_path=None, - output_format=output_format, - meta=meta, - agent_path=agent_path, - skip_reason=f"missing_required_{source_type}", - day=day, - segment=segment, - span_mode=span_mode, - ) + config["skip_reason"] = f"missing_required_{source_type}" + return # Skip when there's nothing to analyze if total_count == 0 or len(transcript.strip()) < MIN_INPUT_CHARS: - return InputContext( - transcript=transcript, - source_counts=source_counts, - prompt=user_prompt, - system_instruction=system_instruction, - system_prompt_name=system_prompt_name, - output_path=None, - output_format=output_format, - meta=meta, - agent_path=agent_path, - skip_reason="no_input", - day=day, - segment=segment, - span_mode=span_mode, - ) + config["skip_reason"] = "no_input" + return # Prepend input context note for limited recordings if total_count < 3: @@ -700,7 +424,7 @@ def assemble_inputs(config: dict) -> InputContext: "**Input Note:** Limited recordings for this day. " "Scale analysis to available input.\n\n" ) - transcript = input_note + transcript + config["transcript"] = input_note + transcript # Build context for template substitution prompt_context: dict[str, str] = {} @@ -743,299 +467,316 @@ def assemble_inputs(config: dict) -> InputContext: agent_path.stem, base_dir=agent_path.parent, context=prompt_context ) prompt = agent_prompt_obj.text - else: - prompt = user_prompt - - # Append user prompt if both agent prompt and user prompt exist - if agent_path and user_prompt and prompt != user_prompt: - prompt = f"{prompt}\n\n{user_prompt}" + # Append user prompt if both exist + if user_prompt and prompt != user_prompt: + prompt = f"{prompt}\n\n{user_prompt}" + config["prompt"] = prompt # Determine output path - output_path: Optional[Path] = None if output_format: if output_path_override: - output_path = Path(output_path_override) + config["output_path"] = Path(output_path_override) elif day: day_dir = str(day_path(day)) - output_path = get_output_path( + config["output_path"] = get_output_path( day_dir, name, segment=segment, output_format=output_format, facet=facet ) - return InputContext( - transcript=transcript, - source_counts=source_counts, - prompt=prompt, - system_instruction=system_instruction, - system_prompt_name=system_prompt_name, - output_path=output_path, - output_format=output_format, - meta=meta, - agent_path=agent_path, - skip_reason=None, - day=day, - segment=segment, - span_mode=span_mode, - ) - # ============================================================================= -# Unified Execution Helpers (shared by generate and tools paths) +# Hook Execution # ============================================================================= -def _emit_start_event( - emit_event: Callable[[dict], None], - name: str, - model: str, - provider: str, - prompt: str, - continue_from: Optional[str] = None, -) -> None: - """Emit a unified start event for both generate and tools paths.""" - start_event: dict[str, Any] = { - "event": "start", - "ts": now_ms(), - "prompt": prompt, - "name": name, - "model": model or "unknown", - "provider": provider, - } - if continue_from: - start_event["continue_from"] = continue_from - emit_event(start_event) - - -def _handle_skip( - inputs: InputContext, - name: str, - path_type: str, - emit_event: Callable[[dict], None], -) -> bool: - """Handle skip conditions from input assembly. +def _run_pre_hooks(config: dict) -> dict: + """Run pre-processing hooks, return dict of modifications. Args: - inputs: InputContext with skip_reason set - name: Agent/generator name - path_type: "generate" or "agent" for logging - emit_event: Event emitter callback + config: Full config dict (hooks receive this directly) Returns: - True if skipped and caller should return/continue, False otherwise + Dict of field modifications to apply to config """ - if not inputs.skip_reason: - return False + meta = config.get("meta", {}) + pre_hook = load_pre_hook(meta) + if not pre_hook: + return {} - logging.info("Config %s skipped: %s", name, inputs.skip_reason) - emit_event( - { - "event": "finish", - "ts": now_ms(), - "result": "", - "skipped": inputs.skip_reason, - } - ) - if inputs.day: - day_log(inputs.day, f"{path_type} {name} skipped ({inputs.skip_reason})") - return True + try: + modifications = pre_hook(config) + if modifications: + LOG.info("Pre-hook returned modifications: %s", list(modifications.keys())) + return modifications + except Exception as exc: + LOG.error("Pre-hook failed: %s", exc) + + return {} -def _execute_pre_hooks( - meta: dict, - modifiable: dict[str, str], - output_path: Optional[Path] = None, - day: Optional[str] = None, - segment: Optional[str] = None, - span_mode: bool = False, -) -> tuple[dict[str, str], dict[str, Any]]: - """Execute pre-processing hooks and return modified values. +def _run_post_hooks(result: str, config: dict) -> str: + """Run post-processing hooks, return transformed result. Args: - meta: Agent metadata containing hook config - modifiable: Dict of modifiable field values (prompt, system_instruction, - transcript, etc.) - these are passed to the hook context - output_path: Output path for context - day: Day string for context - segment: Segment string for context - span_mode: Whether in span mode + result: LLM output text + config: Full config dict (hooks receive this directly) Returns: - Tuple of (modified_values, hook_info for dry-run) - modified_values contains only fields that were modified - hook_info contains name and list of modifications for dry-run display + Transformed result (or original if no hook) """ - hook_info: dict[str, Any] = {} - modified: dict[str, str] = {} + meta = config.get("meta", {}) + post_hook = load_post_hook(meta) + if not post_hook: + return result - pre_hook = load_pre_hook(meta) - if not pre_hook: - return modified, hook_info - - # Get hook name for logging - hook_config = meta.get("hook", {}) - hook_name = hook_config.get("pre") if isinstance(hook_config, dict) else None - hook_info["name"] = hook_name - - # Build context with all modifiable fields - # Note: modifiable may contain transcript, so we don't pass it separately - pre_context = build_pre_hook_context( - meta, - output_path=str(output_path) if output_path else "", - day=day, - segment=segment, - span=span_mode, - **modifiable, - ) + try: + hook_result = post_hook(result, config) + if hook_result is not None: + LOG.info("Post-hook transformed result") + return hook_result + except Exception as exc: + LOG.error("Post-hook failed: %s", exc) - modifications = run_pre_hook(pre_context, pre_hook) - if modifications: - # Only include fields that were actually modified - for key in modifiable: - if key in modifications: - modified[key] = modifications[key] - hook_info["modifications"] = list(modifications.keys()) + return result - return modified, hook_info +# ============================================================================= +# Unified Agent Execution +# ============================================================================= -def _build_dry_run_event( - run_type: str, - name: str, - provider: str, - model: str, - config: dict, - inputs: Optional[InputContext], - hook_info: dict[str, Any], - before_values: dict[str, str], - current_values: dict[str, str], -) -> dict[str, Any]: - """Build a dry-run event with all context. - Args: - run_type: "generate" or "agent" - name: Agent/generator name - provider: Provider name - model: Model name - config: Full config dict - inputs: InputContext if available - hook_info: Pre-hook info from _execute_pre_hooks - before_values: Values before hook execution - current_values: Values after hook execution +def _write_output(output_path: Path, result: str) -> None: + """Write result to output file.""" + output_path.parent.mkdir(parents=True, exist_ok=True) + with open(output_path, "w", encoding="utf-8") as f: + f.write(result) + LOG.info("Wrote output to %s", output_path) + + +def _build_dry_run_event(config: dict, before_values: dict) -> dict: + """Build a dry-run event with all context.""" + has_tools = bool(config.get("tools")) + run_type = "agent" if has_tools else "generate" - Returns: - Complete dry-run event dict - """ event: dict[str, Any] = { "event": "dry_run", "ts": now_ms(), "type": run_type, - "name": name, - "provider": provider, - "model": model or "unknown", - "system_instruction": current_values.get("system_instruction", ""), - "prompt": current_values.get("prompt", ""), + "name": config.get("name", "default"), + "provider": config.get("provider", ""), + "model": config.get("model") or "unknown", + "system_instruction": config.get("system_instruction", ""), + "prompt": config.get("prompt", ""), } - # Add agent-specific fields - if run_type == "agent": - event["user_instruction"] = current_values.get("user_instruction", "") - event["extra_context"] = current_values.get("extra_context", "") + if has_tools: + event["user_instruction"] = config.get("user_instruction", "") + event["extra_context"] = config.get("extra_context", "") event["tools"] = config.get("tools", []) - - # Add generate-specific fields - if run_type == "generate": - event["system_instruction_source"] = ( - inputs.system_prompt_name if inputs else "journal" - ) - event["prompt_source"] = ( - str(inputs.agent_path) if inputs and inputs.agent_path else "request" - ) - - # Add day-based fields if inputs available - if inputs: - event["day"] = inputs.day - event["segment"] = inputs.segment - transcript = current_values.get("transcript", "") + else: + event["system_instruction_source"] = config.get("system_prompt_name", "journal") + agent_path = config.get("agent_path") + event["prompt_source"] = str(agent_path) if agent_path else "request" + + # Day-based fields + if config.get("day"): + event["day"] = config["day"] + event["segment"] = config.get("segment") + transcript = config.get("transcript", "") if transcript: event["transcript"] = transcript event["transcript_chars"] = len(transcript) - event["transcript_files"] = sum(inputs.source_counts.values()) - if inputs.output_path: - event["output_path"] = str(inputs.output_path) - - # Add hook before/after info - if hook_info: - event["pre_hook"] = hook_info.get("name") - event["pre_hook_modifications"] = hook_info.get("modifications", []) - # Include before values for modified fields - for key, before_val in before_values.items(): - current_val = current_values.get(key, "") - if current_val != before_val: - if key == "transcript": - event["transcript_before"] = before_val - event["transcript_before_chars"] = len(before_val) - else: - event[f"{key}_before"] = before_val + event["transcript_files"] = sum(config.get("source_counts", {}).values()) + output_path = config.get("output_path") + if output_path: + event["output_path"] = str(output_path) + + # Show before values for comparison + for key, before_val in before_values.items(): + current_val = config.get(key, "") + if current_val != before_val: + if key == "transcript": + event["transcript_before_chars"] = len(before_val) + else: + event[f"{key}_before"] = before_val return event -def _write_output(output_path: Path, result: str, output_format: str) -> None: - """Write result to output file. +async def _run_agent( + config: dict, + emit_event: Callable[[dict], None], + dry_run: bool = False, +) -> None: + """Execute agent or generator based on config. + + Unified execution path for both tool-using agents and transcript generators. + The only branch is at the LLM call - everything else is shared. Args: - output_path: Path to write to - result: Content to write - output_format: Format type ('md' or 'json') + config: Fully hydrated and enriched config dict + emit_event: Callback to emit JSONL events + dry_run: If True, emit dry_run event instead of calling LLM """ - output_path.parent.mkdir(parents=True, exist_ok=True) - with open(output_path, "w", encoding="utf-8") as f: - f.write(result) - logging.info("Wrote output to %s", output_path) + name = config.get("name", "default") + provider = config.get("provider", "google") + model = config.get("model") + has_tools = bool(config.get("tools")) + force = config.get("force", False) + # Emit start event + start_event: dict[str, Any] = { + "event": "start", + "ts": now_ms(), + "prompt": config.get("prompt", ""), + "name": name, + "model": model or "unknown", + "provider": provider, + } + if config.get("continue_from"): + start_event["continue_from"] = config["continue_from"] + emit_event(start_event) -def _execute_post_hooks( - result: str, - meta: dict, - transcript: str = "", - output_path: Optional[Path] = None, - day: Optional[str] = None, - segment: Optional[str] = None, - span_mode: bool = False, - name: str = "", -) -> str: - """Execute post-processing hooks and return transformed result. + # Handle skip conditions + skip_reason = config.get("skip_reason") + if skip_reason: + LOG.info("Config %s skipped: %s", name, skip_reason) + emit_event( + { + "event": "finish", + "ts": now_ms(), + "result": "", + "skipped": skip_reason, + } + ) + if config.get("day"): + day_log(config["day"], f"agent {name} skipped ({skip_reason})") + return - Args: - result: LLM output text - meta: Agent metadata containing hook config - transcript: Transcript for context - output_path: Output path for context - day: Day for context - segment: Segment for context - span_mode: Span mode flag - name: Agent name + # Check if output already exists (generators only, not tool agents) + output_path = config.get("output_path") + output_format = config.get("output") + if not has_tools and output_path and not force and not dry_run: + if output_path.exists() and output_path.stat().st_size > 0: + LOG.info("Output exists, loading: %s", output_path) + with open(output_path, "r") as f: + result = f.read() + emit_event( + { + "event": "finish", + "ts": now_ms(), + "result": result, + } + ) + return - Returns: - Transformed result (or original if no hook or hook returns None) - """ - post_hook = load_post_hook(meta) - if not post_hook: - return result + # Capture state before pre-hooks + before_values = { + "prompt": config.get("prompt", ""), + "system_instruction": config.get("system_instruction", ""), + "transcript": config.get("transcript", ""), + } + if has_tools: + before_values["user_instruction"] = config.get("user_instruction", "") + before_values["extra_context"] = config.get("extra_context", "") + + # Run pre-hooks + modifications = _run_pre_hooks(config) + for key, value in modifications.items(): + config[key] = value + + # Dry-run mode + if dry_run: + emit_event(_build_dry_run_event(config, before_values)) + return - hook_context = build_hook_context( - meta, - name=name, - day=day, - segment=segment, - span=span_mode, - output_path=str(output_path) if output_path else "", - transcript=transcript, - ) - return run_post_hook(result, hook_context, post_hook) + # Execute LLM call - this is the only real branch + if has_tools: + # Tool-using agent path + from .providers import PROVIDER_REGISTRY, get_provider_module + + if provider not in PROVIDER_REGISTRY: + valid = ", ".join(sorted(PROVIDER_REGISTRY.keys())) + raise ValueError( + f"Unknown provider: {provider!r}. Valid providers: {valid}" + ) + + provider_mod = get_provider_module(provider) + + # Create wrapper to intercept finish event + def agent_emit_event(data: Event) -> None: + if data.get("event") == "finish": + result = data.get("result", "") + result = _run_post_hooks(result, config) + if result != data.get("result", ""): + data = {**data, "result": result} + if output_path and result: + _write_output(output_path, result) + if config.get("handoff"): + data = {**data, "handoff": config["handoff"]} + + # Filter out start events from providers (we already emitted ours) + if data.get("event") == "start": + return + + emit_event(data) + + await provider_mod.run_tools(config=config, on_event=agent_emit_event) + + else: + # Generator path - single-shot generation + from think.models import generate_with_result + from think.muse import key_to_context + + transcript = config.get("transcript", "") + prompt = config.get("prompt", "") + system_instruction = config.get("system_instruction", "") + meta = config.get("meta", {}) + + # Get generation parameters + thinking_budget = meta.get("thinking_budget") or 8192 * 3 + max_output_tokens = meta.get("max_output_tokens") or 8192 * 6 + is_json_output = output_format == "json" + + context = key_to_context(name) + gen_result = generate_with_result( + contents=[transcript, prompt] if transcript else [prompt], + context=context, + temperature=0.3, + max_output_tokens=max_output_tokens, + thinking_budget=thinking_budget, + system_instruction=system_instruction, + json_output=is_json_output, + ) + + result = gen_result["text"] + usage_data = gen_result.get("usage") + + # Run post-hooks + result = _run_post_hooks(result, config) + + # Write output + if output_path and result: + _write_output(output_path, result) + + # Emit finish event + finish_event: dict[str, Any] = { + "event": "finish", + "ts": now_ms(), + "result": result, + } + if usage_data: + finish_event["usage"] = usage_data + if config.get("handoff"): + finish_event["handoff"] = config["handoff"] + emit_event(finish_event) + + # Log completion + if config.get("day"): + day_log(config["day"], f"agent {name} ok") # ============================================================================= -# Generator Functions (for transcript analysis without tools) +# Utility Functions # ============================================================================= @@ -1044,10 +785,6 @@ def scan_day(day: str) -> dict[str, list[str]]: Only scans daily generators (schedule='daily'). Segment generators are stored within segment directories and are not included here. - - Note: Multi-facet generators would produce {topic}_{facet}.{ext} files, - but currently no multi-facet generators exist (all multi-facet agents - have tools, not output). """ day_dir = day_path(day) daily_generators = get_muse_configs( @@ -1065,247 +802,13 @@ def scan_day(day: str) -> dict[str, list[str]]: return {"processed": sorted(processed), "repairable": sorted(pending)} -def generate_agent_output( - transcript: str, - prompt: str, - name: str | None = None, - json_output: bool = False, - system_instruction: str | None = None, - thinking_budget: int | None = None, - max_output_tokens: int | None = None, - return_result: bool = False, -) -> str | GenerateResult: - """Send clustered transcript to LLM for agent output generation. - - Args: - transcript: Clustered transcript content (markdown format). - prompt: Agent prompt text. - name: Agent name for token logging context. - json_output: If True, request JSON response format. - system_instruction: System instruction text. If None, loads default - from journal.md via compose_instructions(). - thinking_budget: Token budget for model thinking. If None, uses default. - max_output_tokens: Maximum output tokens. If None, uses default. - return_result: If True, return full GenerateResult with usage data. - - Returns: - Generated agent output content (markdown or JSON string), or - GenerateResult dict if return_result=True. - """ - from think.models import generate_with_result - - # Use provided system_instruction or fall back to default - if system_instruction is None: - instructions = compose_instructions(include_datetime=False) - system_instruction = instructions["system_instruction"] - - # Use defaults if not specified - if thinking_budget is None: - thinking_budget = 8192 * 3 - if max_output_tokens is None: - max_output_tokens = 8192 * 6 - - # Build context for provider routing and token logging - from think.muse import key_to_context - - context = key_to_context(name) if name else "muse.system.unknown" - - result = generate_with_result( - contents=[transcript, prompt], - context=context, - temperature=0.3, - max_output_tokens=max_output_tokens, - thinking_budget=thinking_budget, - system_instruction=system_instruction, - json_output=json_output, - ) - - if return_result: - return result - return result["text"] - - -def _run_generate( - config: dict, - emit_event: Callable[[dict], None], - *, - dry_run: bool = False, - inputs: InputContext | None = None, -) -> None: - """Execute single-shot generation with optional features based on config. - - This is the generation path for non-tool requests. Uses assemble_inputs() for - input assembly, which can be shared with the tools path. - - Args: - config: Merged config from cortex - emit_event: Callback to emit JSONL events - dry_run: If True, emit dry_run event instead of calling LLM - inputs: Pre-assembled InputContext (if None, will call assemble_inputs) - """ - name = config.get("name", "default") - force = config.get("force", False) - provider = config.get("provider", "google") - model = config.get("model") - user_prompt = config.get("prompt", "") - - # Assemble inputs first (before start event, so we can skip cleanly) - if inputs is None: - inputs = assemble_inputs(config) - - # Emit unified start event - _emit_start_event(emit_event, name, model, provider, user_prompt) - - # Handle skip conditions using helper - if _handle_skip(inputs, name, "generate", emit_event): - return - - # Extract values from InputContext - day = inputs.day - segment = inputs.segment - span_mode = inputs.span_mode - transcript = inputs.transcript - prompt = inputs.prompt - system_instruction = inputs.system_instruction - output_path = inputs.output_path - output_format = inputs.output_format - meta = inputs.meta - - # Check if output exists - output_exists = False - is_json_output = output_format == "json" - if output_path: - output_exists = output_path.exists() and output_path.stat().st_size > 0 - - # Extract generation parameters from metadata - meta_thinking_budget = meta.get("thinking_budget") - meta_max_output_tokens = meta.get("max_output_tokens") - - usage_data = None - - # Dry-run always goes through prompt assembly, regardless of existing output - if output_exists and not force and not dry_run: - # Load existing content (no LLM call) - logging.info("Output exists, loading: %s", output_path) - with open(output_path, "r") as f: - result = f.read() - else: - # Generate new content - if output_exists and force: - logging.info("Force regenerating: %s", output_path) - - # Capture state before pre-hook for dry-run comparison - before_values = { - "transcript": transcript, - "prompt": prompt, - "system_instruction": system_instruction, - } - - # Run pre-processing hooks using helper - modifications, hook_info = _execute_pre_hooks( - meta, - modifiable={ - "prompt": prompt, - "system_instruction": system_instruction, - "transcript": transcript, - }, - output_path=output_path, - day=day, - segment=segment, - span_mode=span_mode, - ) - - # Apply modifications - transcript = modifications.get("transcript", transcript) - prompt = modifications.get("prompt", prompt) - system_instruction = modifications.get("system_instruction", system_instruction) - - # Current values after hook - current_values = { - "transcript": transcript, - "prompt": prompt, - "system_instruction": system_instruction, - } - - # Dry-run mode: emit context and return without LLM call - if dry_run: - dry_run_event = _build_dry_run_event( - "generate", - name, - provider, - model, - config, - inputs, - hook_info, - before_values, - current_values, - ) - emit_event(dry_run_event) - return - - gen_result = generate_agent_output( - transcript, - prompt, - name=name, - json_output=is_json_output, - system_instruction=system_instruction, - thinking_budget=meta_thinking_budget, - max_output_tokens=meta_max_output_tokens, - return_result=True, - ) - result = gen_result["text"] - usage_data = gen_result.get("usage") - - # Run post-processing hooks using helper - result = _execute_post_hooks( - result, - meta, - transcript=transcript, - output_path=output_path, - day=day, - segment=segment, - span_mode=span_mode, - name=name, - ) - - # Write output file (agents.py owns output writing) - if output_path and result: - _write_output(output_path, result, output_format or "md") - - # Emit finish event with result - finish_event: dict[str, Any] = { - "event": "finish", - "ts": now_ms(), - "result": result, - } - if usage_data: - finish_event["usage"] = usage_data - # Include handoff config for cortex to spawn follow-up agent - if config.get("handoff"): - finish_event["handoff"] = config["handoff"] - - emit_event(finish_event) - - # Log completion (only for day-based requests) - if day: - msg = f"generate {name} ok" - if force: - msg += " --force" - day_log(day, msg) - - # ============================================================================= # Main Entry Point # ============================================================================= async def main_async() -> None: - """NDJSON-based CLI for agents and generators. - - Routes based on config: - - 'output' field present (no 'tools') -> generator (transcript analysis) - - Everything else -> agent (with or without tools, via provider) - """ + """NDJSON-based CLI for agents and generators.""" parser = argparse.ArgumentParser( description="solstone Agent CLI - Accepts NDJSON input via stdin" ) @@ -1319,8 +822,6 @@ async def main_async() -> None: dry_run = args.dry_run app_logger = setup_logging(args.verbose) - - # Always write to stdout only event_writer = JSONEventWriter(None) def emit_event(data: Event) -> None: @@ -1329,7 +830,6 @@ async def main_async() -> None: event_writer.emit(data) try: - # NDJSON input mode from stdin only app_logger.info("Processing NDJSON input from stdin") for line in sys.stdin: line = line.strip() @@ -1337,189 +837,16 @@ async def main_async() -> None: continue try: - # Parse NDJSON line - raw request from cortex request = json.loads(line) - - # Hydrate config: load agent definition, merge request, resolve provider config = hydrate_config(request) - # Validate config error = validate_config(config) if error: emit_event({"event": "error", "error": error, "ts": now_ms()}) continue - # Route based on config type: tools → run_tools, else → run_generate - has_tools = bool(config.get("tools")) - - if not has_tools: - # Generate path: single-shot generation with opt-in features - app_logger.debug(f"Processing generate: {config.get('name')}") - _run_generate(config, emit_event, dry_run=dry_run) - - else: - # Agent: with or without tools (conversational or tool-using) - # Extract provider to route to correct module - from .providers import PROVIDER_REGISTRY, get_provider_module - - provider = config.get("provider", "google") - name = config.get("name", "default") - model = config.get("model") - - app_logger.debug(f"Processing agent: provider={provider}") - - # Route to appropriate provider module - if provider in PROVIDER_REGISTRY: - provider_mod = get_provider_module(provider) - else: - # Explicit error for unknown providers - valid = ", ".join(sorted(PROVIDER_REGISTRY.keys())) - raise ValueError( - f"Unknown provider: {provider!r}. Valid providers: {valid}" - ) - - # Assemble inputs if day is specified (transcript loading) - inputs: InputContext | None = None - if config.get("day"): - inputs = assemble_inputs(config) - - # Get metadata for hooks (from inputs if available, else from config) - meta = inputs.meta if inputs else config - - # Emit unified start event (agents.py owns this) - _emit_start_event( - emit_event, - name, - model, - provider, - config.get("prompt", ""), - continue_from=config.get("continue_from"), - ) - - # Handle skip conditions using helper - if inputs and _handle_skip(inputs, name, "agent", emit_event): - continue - - # Pass transcript and system instruction to provider if inputs assembled - if inputs and not config.get("continue_from"): - # Warn if both day and continue_from are specified - if config.get("continue_from"): - logging.warning( - "Both 'day' and 'continue_from' specified; " - "continue_from takes precedence, transcript ignored" - ) - else: - config["transcript"] = inputs.transcript - if not config.get("system_instruction"): - config["system_instruction"] = inputs.system_instruction - - # Capture state before pre-hook for dry-run comparison - before_values = { - "prompt": config.get("prompt", ""), - "system_instruction": config.get("system_instruction", ""), - "user_instruction": config.get("user_instruction", ""), - "extra_context": config.get("extra_context", ""), - "transcript": config.get("transcript", ""), - } - - # Run pre-processing hooks using helper - # Note: before_values already contains transcript - modifications, hook_info = _execute_pre_hooks( - meta, - modifiable=before_values.copy(), - output_path=inputs.output_path if inputs else None, - day=inputs.day if inputs else None, - segment=inputs.segment if inputs else None, - span_mode=inputs.span_mode if inputs else False, - ) - - # Apply modifications to config - for key in ( - "prompt", - "system_instruction", - "user_instruction", - "extra_context", - "transcript", - ): - if key in modifications: - config[key] = modifications[key] - - # Current values after hook - current_values = { - "prompt": config.get("prompt", ""), - "system_instruction": config.get("system_instruction", ""), - "user_instruction": config.get("user_instruction", ""), - "extra_context": config.get("extra_context", ""), - "transcript": config.get("transcript", ""), - } - - # Dry-run mode: emit context and return without LLM call - if dry_run: - dry_run_event = _build_dry_run_event( - "agent", - name, - provider, - model, - config, - inputs, - hook_info, - before_values, - current_values, - ) - emit_event(dry_run_event) - continue - - handoff_config = config.get("handoff") - output_path = inputs.output_path if inputs else None - output_format = inputs.output_format if inputs else None - - # Create event handler that intercepts finish for post-hooks, - # output writing, and handoff - def agent_emit_event(data: Event) -> None: - if data.get("event") == "finish": - result = data.get("result", "") - - # Apply post-processing hooks using helper - result = _execute_post_hooks( - result, - meta, - transcript=config.get("transcript", ""), - output_path=output_path, - day=inputs.day if inputs else None, - segment=inputs.segment if inputs else None, - span_mode=inputs.span_mode if inputs else False, - name=name, - ) - - # Update data if result was transformed - if result != data.get("result", ""): - data = {**data, "result": result} - - # Write output file (agents.py owns output writing) - if output_path and result: - _write_output( - output_path, result, output_format or "md" - ) - - # Include handoff config for cortex - if handoff_config: - data = {**data, "handoff": handoff_config} - - # Filter out start events from providers (we already emitted ours) - if data.get("event") == "start": - return - - emit_event(data) - - # Pass complete config to provider - await provider_mod.run_tools( - config=config, - on_event=agent_emit_event, - ) - - # Log completion for day-based requests - if inputs and inputs.day: - day_log(inputs.day, f"agent {name} ok") + enrich_config(config) + await _run_agent(config, emit_event, dry_run=dry_run) except json.JSONDecodeError as e: emit_event( @@ -1539,7 +866,7 @@ async def main_async() -> None: } ) - except Exception as exc: # pragma: no cover - unexpected + except Exception as exc: err = { "event": "error", "error": str(exc), @@ -1554,5 +881,15 @@ async def main_async() -> None: def main() -> None: """Entry point wrapper.""" - asyncio.run(main_async()) + + +__all__ = [ + "format_tool_summary", + "parse_agent_events_to_turns", + "hydrate_config", + "expand_tools", + "validate_config", + "enrich_config", + "scan_day", +] diff --git a/think/muse.py b/think/muse.py index 67c325e4f..504801ab4 100644 --- a/think/muse.py +++ b/think/muse.py @@ -705,6 +705,6 @@ def load_pre_hook(config: dict) -> Callable[["PreHookContext"], dict | None] | N return _load_hook_function(config, "pre", "pre_process") -# Type aliases for hook context (actual TypedDicts defined in agents.py) +# Type aliases for hook context - hooks receive the full config dict HookContext = dict PreHookContext = dict diff --git a/think/providers/anthropic.py b/think/providers/anthropic.py index 19f0df75a..9d17b475c 100644 --- a/think/providers/anthropic.py +++ b/think/providers/anthropic.py @@ -53,7 +53,6 @@ from .shared import ( GenerateResult, JSONEventCallback, ThinkingEvent, - extract_agent_config, extract_tool_result, ) @@ -258,7 +257,21 @@ async def run_tools( user_instruction, extra_context, model, etc. on_event: Optional event callback """ - ac = extract_agent_config(config, default_max_tokens=_DEFAULT_MAX_TOKENS) + # Extract config values directly + prompt = config.get("prompt", "") + model = config.get("model", _DEFAULT_MODEL) + system_instruction = config.get("system_instruction") + user_instruction = config.get("user_instruction") + extra_context = config.get("extra_context") + transcript = config.get("transcript") + mcp_server_url = config.get("mcp_server_url") + tools_filter = config.get("tools") + max_output_tokens = config.get("max_output_tokens", _DEFAULT_MAX_TOKENS) + thinking_budget_config = config.get("thinking_budget") + continue_from = config.get("continue_from") + agent_id = config.get("agent_id") + name = config.get("name") + callback = JSONEventCallback(on_event) try: @@ -271,46 +284,46 @@ async def run_tools( # Note: Start event is emitted by agents.py (unified event ownership) # Build initial messages - check for continuation first - if ac.continue_from: + if continue_from: # Load previous conversation history using shared function from ..agents import parse_agent_events_to_turns - messages = parse_agent_events_to_turns(ac.continue_from) + messages = parse_agent_events_to_turns(continue_from) # Add new prompt as continuation - messages.append({"role": "user", "content": ac.prompt}) + messages.append({"role": "user", "content": prompt}) else: # Fresh conversation messages: list[MessageParam] = [] # Prepend transcript if provided (from day/segment input assembly) - if ac.transcript: - messages.append({"role": "user", "content": ac.transcript}) - if ac.extra_context: - messages.append({"role": "user", "content": ac.extra_context}) - if ac.user_instruction: - messages.append({"role": "user", "content": ac.user_instruction}) - messages.append({"role": "user", "content": ac.prompt}) + if transcript: + messages.append({"role": "user", "content": transcript}) + if extra_context: + messages.append({"role": "user", "content": extra_context}) + if user_instruction: + messages.append({"role": "user", "content": user_instruction}) + messages.append({"role": "user", "content": prompt}) # Initialize tools and executor if MCP server URL provided - if ac.mcp_server_url: - async with create_mcp_client(str(ac.mcp_server_url)) as mcp: - if ac.tools and isinstance(ac.tools, list): - logger.info(f"Using tool filter with allowed tools: {ac.tools}") + if mcp_server_url: + async with create_mcp_client(str(mcp_server_url)) as mcp: + if tools_filter and isinstance(tools_filter, list): + logger.info(f"Using tool filter with allowed tools: {tools_filter}") - tools = await _get_mcp_tools(mcp, ac.tools) + tools = await _get_mcp_tools(mcp, tools_filter) tool_executor = ToolExecutor( - mcp, callback, agent_id=ac.agent_id, name=ac.name + mcp, callback, agent_id=agent_id, name=name ) thinking_budget, effective_max_tokens = _resolve_agent_thinking_params( - ac.max_output_tokens, ac.thinking_budget + max_output_tokens, thinking_budget_config ) for _ in range(_MAX_TOOL_ITERATIONS): # Build request params - thinking always enabled create_params = { - "model": ac.model, + "model": model, "max_tokens": effective_max_tokens, - "system": ac.system_instruction, + "system": system_instruction, "messages": messages, "thinking": { "type": "enabled", @@ -330,7 +343,7 @@ async def run_tools( elif getattr(block, "type", None) == "tool_use": tool_uses.append(block) elif isinstance(block, (ThinkingBlock, RedactedThinkingBlock)): - _emit_thinking_event(block, ac.model, callback) + _emit_thinking_event(block, model, callback) messages.append({"role": "assistant", "content": response.content}) @@ -383,12 +396,12 @@ async def run_tools( else: # No MCP tools - single response only thinking_budget, effective_max_tokens = _resolve_agent_thinking_params( - ac.max_output_tokens, ac.thinking_budget + max_output_tokens, thinking_budget_config ) create_params = { - "model": ac.model, + "model": model, "max_tokens": effective_max_tokens, - "system": ac.system_instruction, + "system": system_instruction, "messages": messages, "thinking": { "type": "enabled", @@ -403,7 +416,7 @@ async def run_tools( if getattr(block, "type", None) == "text": final_text += block.text elif isinstance(block, (ThinkingBlock, RedactedThinkingBlock)): - _emit_thinking_event(block, ac.model, callback) + _emit_thinking_event(block, model, callback) finish_event = { "event": "finish", diff --git a/think/providers/google.py b/think/providers/google.py index 1877f216a..34d100aeb 100644 --- a/think/providers/google.py +++ b/think/providers/google.py @@ -47,7 +47,6 @@ from .shared import ( GenerateResult, JSONEventCallback, ThinkingEvent, - extract_agent_config, extract_tool_result, ) @@ -540,7 +539,21 @@ async def run_tools( user_instruction, extra_context, model, etc. on_event: Optional event callback """ - ac = extract_agent_config(config, default_max_tokens=_DEFAULT_MAX_TOKENS) + # Extract config values directly + prompt = config.get("prompt", "") + model = config.get("model", _DEFAULT_MODEL) + system_instruction = config.get("system_instruction") + user_instruction = config.get("user_instruction") + extra_context = config.get("extra_context") + transcript = config.get("transcript") + mcp_server_url = config.get("mcp_server_url") + tools = config.get("tools") + max_output_tokens = config.get("max_output_tokens", _DEFAULT_MAX_TOKENS) + thinking_budget = config.get("thinking_budget") + continue_from = config.get("continue_from") + agent_id = config.get("agent_id") + name = config.get("name") + callback = JSONEventCallback(on_event) try: @@ -552,11 +565,11 @@ async def run_tools( # Note: Start event is emitted by agents.py (unified event ownership) # Build history - check for continuation first - if ac.continue_from: + if continue_from: # Load previous conversation history using shared function from ..agents import parse_agent_events_to_turns - turns = parse_agent_events_to_turns(ac.continue_from) + turns = parse_agent_events_to_turns(continue_from) # Convert to Google's format history = [] for turn in turns: @@ -568,20 +581,18 @@ async def run_tools( # Fresh conversation - convert generic turns to Google format history = [] # Prepend transcript if provided (from day/segment input assembly) - if ac.transcript: + if transcript: history.append( - types.Content(role="user", parts=[types.Part(text=ac.transcript)]) + types.Content(role="user", parts=[types.Part(text=transcript)]) ) - if ac.extra_context: + if extra_context: history.append( - types.Content( - role="user", parts=[types.Part(text=ac.extra_context)] - ) + types.Content(role="user", parts=[types.Part(text=extra_context)]) ) - if ac.user_instruction: + if user_instruction: history.append( types.Content( - role="user", parts=[types.Part(text=ac.user_instruction)] + role="user", parts=[types.Part(text=user_instruction)] ) ) @@ -593,10 +604,8 @@ async def run_tools( # Create fresh chat session chat = client.aio.chats.create( - model=ac.model, - config=types.GenerateContentConfig( - system_instruction=ac.system_instruction - ), + model=model, + config=types.GenerateContentConfig(system_instruction=system_instruction), history=history, ) @@ -604,29 +613,25 @@ async def run_tools( tool_call_count = 0 # Configure tools if MCP server URL provided - if ac.mcp_server_url: + if mcp_server_url: # Create MCP client and attach hooks - async with create_mcp_client(str(ac.mcp_server_url)) as mcp: + async with create_mcp_client(str(mcp_server_url)) as mcp: # Attach tool logging hooks to the MCP session - tool_hooks = ToolLoggingHooks( - callback, agent_id=ac.agent_id, name=ac.name - ) + tool_hooks = ToolLoggingHooks(callback, agent_id=agent_id, name=name) tool_hooks.attach(mcp.session) # Configure function calling mode based on tool filtering - if ac.tools and isinstance(ac.tools, list): - logger.info(f"Filtering tools to: {ac.tools}") + if tools and isinstance(tools, list): + logger.info(f"Filtering tools to: {tools}") function_calling_config = types.FunctionCallingConfig( mode="ANY", # Restrict to only allowed functions - allowed_function_names=ac.tools, + allowed_function_names=tools, ) else: function_calling_config = types.FunctionCallingConfig(mode="AUTO") total_tokens, effective_thinking_budget = ( - _compute_agent_thinking_params( - ac.max_output_tokens, ac.thinking_budget - ) + _compute_agent_thinking_params(max_output_tokens, thinking_budget) ) cfg = types.GenerateContentConfig( @@ -642,15 +647,15 @@ async def run_tools( ) # Send the message - SDK handles automatic function calling - response = await chat.send_message(ac.prompt, config=cfg) - _emit_thinking_events(response, ac.model, callback) + response = await chat.send_message(prompt, config=cfg) + _emit_thinking_events(response, model, callback) # Capture tool call count from hooks tool_call_count = tool_hooks._counter else: # No MCP tools - just basic config total_tokens, effective_thinking_budget = _compute_agent_thinking_params( - ac.max_output_tokens, ac.thinking_budget + max_output_tokens, thinking_budget ) cfg = types.GenerateContentConfig( @@ -661,8 +666,8 @@ async def run_tools( ), ) - response = await chat.send_message(ac.prompt, config=cfg) - _emit_thinking_events(response, ac.model, callback) + response = await chat.send_message(prompt, config=cfg) + _emit_thinking_events(response, model, callback) # Extract finish reason for diagnostics and user-friendly messages finish_reason = _extract_finish_reason(response) diff --git a/think/providers/openai.py b/think/providers/openai.py index 9420bb6a6..2dacfa073 100644 --- a/think/providers/openai.py +++ b/think/providers/openai.py @@ -80,7 +80,6 @@ from .shared import ( GenerateResult, JSONEventCallback, ThinkingEvent, - extract_agent_config, ) @@ -223,42 +222,54 @@ async def run_tools( user_instruction, extra_context, model, etc. on_event: Optional event callback """ - ac = extract_agent_config(config, default_max_tokens=_DEFAULT_MAX_TOKENS) + # Extract config values directly + prompt = config.get("prompt", "") + model = config.get("model", GPT_5) + system_instruction = config.get("system_instruction") + user_instruction = config.get("user_instruction") + extra_context = config.get("extra_context") + transcript = config.get("transcript") + mcp_server_url = config.get("mcp_server_url") + tools = config.get("tools") + max_output_tokens = config.get("max_output_tokens", _DEFAULT_MAX_TOKENS) + continue_from = config.get("continue_from") + agent_id = config.get("agent_id") + name = config.get("name") max_turns = config.get("max_turns", _DEFAULT_MAX_TURNS) - LOG.info("Running agent with model %s", ac.model) + LOG.info("Running agent with model %s", model) cb = JSONEventCallback(on_event) # Note: Start event is emitted by agents.py (unified event ownership) # Model settings: always enable reasoning with detailed summaries model_settings = ModelSettings( - max_tokens=ac.max_output_tokens, + max_tokens=max_output_tokens, reasoning=_DEFAULT_REASONING, ) # Initialize MCP server mcp_server = None - if ac.mcp_server_url: - http_uri = _normalize_streamable_http_uri(str(ac.mcp_server_url)) + if mcp_server_url: + http_uri = _normalize_streamable_http_uri(str(mcp_server_url)) # Configure tool filter if tools are specified tool_filter = None - if ac.tools and isinstance(ac.tools, list) and ToolFilterStatic: + if tools and isinstance(tools, list) and ToolFilterStatic: # Create a tool filter with allowed tools - tool_filter = ToolFilterStatic(allowed_tool_names=ac.tools) - LOG.info(f"Using tool filter with allowed tools: {ac.tools}") - elif ac.tools: + tool_filter = ToolFilterStatic(allowed_tool_names=tools) + LOG.info(f"Using tool filter with allowed tools: {tools}") + elif tools: LOG.warning( "Tool filtering requested but ToolFilterStatic not available in this version" ) mcp_params = {"url": http_uri} headers: dict[str, str] = {} - if ac.agent_id: - headers["X-Agent-Id"] = str(ac.agent_id) - if ac.name: - headers["X-Agent-Name"] = ac.name + if agent_id: + headers["X-Agent-Id"] = str(agent_id) + if name: + headers["X-Agent-Name"] = name if headers: mcp_params["headers"] = headers @@ -280,14 +291,14 @@ async def run_tools( finish_reason_holder = [None] # Create session and load history if continuing conversation - session_id = ac.continue_from or ac.agent_id or f"session-{int(time.time())}" + session_id = continue_from or agent_id or f"session-{int(time.time())}" session = SQLiteSession(session_id=session_id, db_path=":memory:") # Load conversation history if continuing - if ac.continue_from: + if continue_from: from ..agents import parse_agent_events_to_turns - turns = parse_agent_events_to_turns(ac.continue_from) + turns = parse_agent_events_to_turns(continue_from) if turns: items = _convert_turns_to_items(turns) await session.add_items(items) @@ -295,12 +306,12 @@ async def run_tools( # Fresh conversation - add transcript, context and user instruction as initial messages initial_turns = [] # Prepend transcript if provided (from day/segment input assembly) - if ac.transcript: - initial_turns.append({"role": "user", "content": ac.transcript}) - if ac.extra_context: - initial_turns.append({"role": "user", "content": ac.extra_context}) - if ac.user_instruction: - initial_turns.append({"role": "user", "content": ac.user_instruction}) + if transcript: + initial_turns.append({"role": "user", "content": transcript}) + if extra_context: + initial_turns.append({"role": "user", "content": extra_context}) + if user_instruction: + initial_turns.append({"role": "user", "content": user_instruction}) if initial_turns: initial_items = _convert_turns_to_items(initial_turns) await session.add_items(initial_items) @@ -324,15 +335,15 @@ async def run_tools( mcp_servers_list = [mcp_server] if mcp_server else [] agent = Agent( name="solstoneCLI", - instructions=ac.system_instruction, - model=ac.model, + instructions=system_instruction, + model=model, model_settings=model_settings, mcp_servers=mcp_servers_list, ) result = Runner.run_streamed( agent, - input=ac.prompt, + input=prompt, session=session, run_config=RunConfig(tracing_disabled=True), # per docs max_turns=max_turns, @@ -438,7 +449,7 @@ async def run_tools( thinking_event: ThinkingEvent = { "event": "thinking", "summary": summary_text, - "model": ac.model, + "model": model, "ts": now_ms(), } cb.emit(thinking_event) diff --git a/think/providers/shared.py b/think/providers/shared.py index fcd92acd7..4789103b1 100644 --- a/think/providers/shared.py +++ b/think/providers/shared.py @@ -12,7 +12,6 @@ This module contains: from __future__ import annotations -from dataclasses import dataclass from typing import Any, Callable, Literal, Optional, Union from typing_extensions import Required, TypedDict @@ -151,83 +150,6 @@ class JSONEventCallback: pass -# --------------------------------------------------------------------------- -# Agent Config Extraction -# --------------------------------------------------------------------------- - - -@dataclass -class AgentConfig: - """Validated agent configuration extracted from config dict.""" - - prompt: str - model: str - name: str - agent_id: Optional[str] - max_output_tokens: int - thinking_budget: Optional[int] - mcp_server_url: Optional[str] - continue_from: Optional[str] - system_instruction: str - extra_context: str - user_instruction: str - tools: Optional[list[str]] - provider: str - - # Transcript from journal (if day/segment specified) - transcript: str - - # Original config for provider-specific access - raw_config: dict - - -def extract_agent_config(config: dict, default_max_tokens: int = 8192) -> AgentConfig: - """Extract and validate agent configuration. - - Parameters - ---------- - config - Raw config dict from cortex. - default_max_tokens - Default max_output_tokens if not specified in config. - - Returns - ------- - AgentConfig - Validated configuration dataclass. - - Raises - ------ - ValueError - If required fields are missing. - """ - prompt = config.get("prompt", "") - if not prompt: - raise ValueError("Missing 'prompt' in config") - - model = config.get("model") - if not model: - raise ValueError("Missing 'model' in config - should be set by Cortex") - - return AgentConfig( - prompt=prompt, - model=model, - name=config.get("name", "default"), - agent_id=config.get("agent_id"), - max_output_tokens=config.get("max_output_tokens", default_max_tokens), - thinking_budget=config.get("thinking_budget"), - mcp_server_url=config.get("mcp_server_url"), - continue_from=config.get("continue_from"), - system_instruction=config.get("system_instruction", ""), - extra_context=config.get("extra_context", ""), - user_instruction=config.get("user_instruction", ""), - tools=config.get("tools"), - provider=config.get("provider", "google"), - transcript=config.get("transcript", ""), - raw_config=config, - ) - - # --------------------------------------------------------------------------- # MCP Tool Result Extraction # --------------------------------------------------------------------------- @@ -275,6 +197,5 @@ __all__ = [ "GenerateResult", "JSONEventCallback", "ThinkingEvent", - "extract_agent_config", "extract_tool_result", ]