Address cache probe and Gemini review feedback
This commit is contained in:
parent
1b32a8a75c
commit
4ceec9f778
|
|
@ -29,7 +29,7 @@ import json
|
|||
import sys
|
||||
import time
|
||||
|
||||
payload = json.loads(sys.argv[1])
|
||||
payload = json.loads(sys.stdin.read())
|
||||
|
||||
from src.utils.clients import honcho_llm_call_inner
|
||||
|
||||
|
|
@ -66,6 +66,8 @@ async def run() -> None:
|
|||
asyncio.run(run())
|
||||
"""
|
||||
|
||||
PROBE_SUBPROCESS_TIMEOUT_SECONDS = 300
|
||||
|
||||
BASE_PREFIX = "\n".join(
|
||||
[
|
||||
"You are Honcho's memory-backed reasoning layer.",
|
||||
|
|
@ -119,7 +121,13 @@ ROLLING_HISTORY_B = "\n".join(
|
|||
|
||||
@dataclass
|
||||
class ProviderSpec:
|
||||
"""Provider label and model configuration for one comparison run."""
|
||||
"""Provider configuration for one probe run.
|
||||
|
||||
Args:
|
||||
label: Human-readable label for reports.
|
||||
provider: Provider identifier for `honcho_llm_call_inner`.
|
||||
model: Model name to evaluate.
|
||||
"""
|
||||
|
||||
label: str
|
||||
provider: str
|
||||
|
|
@ -128,14 +136,23 @@ class ProviderSpec:
|
|||
|
||||
@dataclass
|
||||
class VariantSpec:
|
||||
"""A named Honcho worktree variant used in the comparison."""
|
||||
"""Named worktree variant used in the comparison.
|
||||
|
||||
Args:
|
||||
label: Short label for reports.
|
||||
worktree: Path to the Honcho worktree to execute in.
|
||||
"""
|
||||
|
||||
label: str
|
||||
worktree: Path
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
"""Parse command-line arguments for the cache comparison probe."""
|
||||
"""Parse probe command-line arguments.
|
||||
|
||||
Returns:
|
||||
Parsed command-line arguments.
|
||||
"""
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Compare prompt-prefix cache behavior between two Honcho worktrees."
|
||||
)
|
||||
|
|
@ -182,7 +199,17 @@ def parse_args() -> argparse.Namespace:
|
|||
|
||||
|
||||
def parse_provider_spec(raw: str) -> ProviderSpec:
|
||||
"""Parse a provider specification of the form label=provider:model."""
|
||||
"""Parse ``label=provider:model`` into a provider spec.
|
||||
|
||||
Args:
|
||||
raw: Raw provider specification.
|
||||
|
||||
Returns:
|
||||
Parsed provider configuration.
|
||||
|
||||
Raises:
|
||||
ValueError: If the provider specification does not match the expected format.
|
||||
"""
|
||||
if "=" not in raw or ":" not in raw:
|
||||
raise ValueError(
|
||||
f"Invalid provider spec {raw!r}. Expected label=provider:model"
|
||||
|
|
@ -193,14 +220,32 @@ def parse_provider_spec(raw: str) -> ProviderSpec:
|
|||
|
||||
|
||||
def add_namespace(namespace: str, content: str) -> str:
|
||||
"""Prefix prompt content with a cache namespace tag."""
|
||||
"""Prefix prompt content with a cache namespace tag.
|
||||
|
||||
Args:
|
||||
namespace: Cache namespace to prepend.
|
||||
content: Prompt content to tag.
|
||||
|
||||
Returns:
|
||||
Tagged prompt content.
|
||||
"""
|
||||
return f"<cache_namespace>{namespace}</cache_namespace>\n{content}"
|
||||
|
||||
|
||||
def build_messages(
|
||||
namespace: str, base_prefix: str, rolling_history: str, user_query: str
|
||||
) -> list[dict[str, str]]:
|
||||
"""Build the system and user messages for a single probe call."""
|
||||
"""Build messages for one probe call.
|
||||
|
||||
Args:
|
||||
namespace: Cache namespace for this call.
|
||||
base_prefix: Stable system prompt content.
|
||||
rolling_history: Dynamic history block.
|
||||
user_query: User query for the call.
|
||||
|
||||
Returns:
|
||||
Provider-agnostic chat messages.
|
||||
"""
|
||||
return [
|
||||
{"role": "system", "content": add_namespace(namespace, base_prefix)},
|
||||
{"role": "system", "content": add_namespace(namespace, rolling_history)},
|
||||
|
|
@ -209,7 +254,14 @@ def build_messages(
|
|||
|
||||
|
||||
def build_scenarios(selected: set[str] | None) -> list[dict[str, Any]]:
|
||||
"""Return the scenario definitions, optionally filtered by name."""
|
||||
"""Return scenario definitions, optionally filtered by name.
|
||||
|
||||
Args:
|
||||
selected: Optional scenario names to include.
|
||||
|
||||
Returns:
|
||||
Scenario definitions in execution order.
|
||||
"""
|
||||
scenario_defs = [
|
||||
{
|
||||
"name": "repeat_exact",
|
||||
|
|
@ -270,17 +322,17 @@ def build_scenarios(selected: set[str] | None) -> list[dict[str, Any]]:
|
|||
"prime": {
|
||||
"base_prefix": BASE_PREFIX,
|
||||
"rolling_history": ROLLING_HISTORY_A,
|
||||
"user_query": "What city is the user considering for a move?",
|
||||
"user_query": "What exact launch date should be remembered for the user?",
|
||||
},
|
||||
"transition": {
|
||||
"base_prefix": BASE_PREFIX_VARIANT,
|
||||
"rolling_history": ROLLING_HISTORY_A,
|
||||
"user_query": "What city is the user considering for a move?",
|
||||
"user_query": "What exact launch date should be remembered for the user?",
|
||||
},
|
||||
"steady": {
|
||||
"base_prefix": BASE_PREFIX_VARIANT,
|
||||
"rolling_history": ROLLING_HISTORY_A,
|
||||
"user_query": "What city is the user considering for a move?",
|
||||
"user_query": "What exact launch date should be remembered for the user?",
|
||||
},
|
||||
},
|
||||
]
|
||||
|
|
@ -295,10 +347,18 @@ def build_cache_namespace(
|
|||
scenario_name: str,
|
||||
invocation_salt: str,
|
||||
) -> str:
|
||||
"""Build a cache namespace that is stable within one run and unique across runs."""
|
||||
return (
|
||||
f"{variant.label}:{provider.label}:{scenario_name}:{invocation_salt}"
|
||||
)
|
||||
"""Build a cache namespace for one invocation.
|
||||
|
||||
Args:
|
||||
variant: Worktree variant under test.
|
||||
provider: Provider/model pair being evaluated.
|
||||
scenario_name: Scenario identifier.
|
||||
invocation_salt: Per-process salt generated once at startup.
|
||||
|
||||
Returns:
|
||||
Namespace stable within one run and unique across runs.
|
||||
"""
|
||||
return f"{variant.label}:{provider.label}:{scenario_name}:{invocation_salt}"
|
||||
|
||||
|
||||
def build_scenario_payload(
|
||||
|
|
@ -308,7 +368,18 @@ def build_scenario_payload(
|
|||
max_tokens: int,
|
||||
invocation_salt: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the subprocess payload for one scenario probe."""
|
||||
"""Build the child-process payload for one scenario.
|
||||
|
||||
Args:
|
||||
variant: Worktree variant under test.
|
||||
provider: Provider/model pair being evaluated.
|
||||
scenario: Scenario definition to execute.
|
||||
max_tokens: Maximum output tokens to request.
|
||||
invocation_salt: Per-process salt generated once at startup.
|
||||
|
||||
Returns:
|
||||
Serialized probe payload.
|
||||
"""
|
||||
namespace = build_cache_namespace(
|
||||
variant,
|
||||
provider,
|
||||
|
|
@ -345,7 +416,21 @@ def run_scenario_probe(
|
|||
max_tokens: int,
|
||||
invocation_salt: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Execute one scenario probe in a subprocess for the given worktree."""
|
||||
"""Run one scenario probe in a subprocess.
|
||||
|
||||
Args:
|
||||
variant: Worktree variant under test.
|
||||
provider: Provider/model pair being evaluated.
|
||||
scenario: Scenario definition to execute.
|
||||
max_tokens: Maximum output tokens to request.
|
||||
invocation_salt: Per-process salt generated once at startup.
|
||||
|
||||
Returns:
|
||||
Parsed JSON result from the child process.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the child times out, fails, or emits invalid JSON.
|
||||
"""
|
||||
payload = build_scenario_payload(
|
||||
variant,
|
||||
provider,
|
||||
|
|
@ -353,13 +438,25 @@ def run_scenario_probe(
|
|||
max_tokens,
|
||||
invocation_salt,
|
||||
)
|
||||
process = subprocess.run(
|
||||
[sys.executable, "-c", CHILD_CODE, json.dumps(payload)],
|
||||
cwd=variant.worktree,
|
||||
text=True,
|
||||
capture_output=True,
|
||||
check=False,
|
||||
)
|
||||
try:
|
||||
process = subprocess.run(
|
||||
[sys.executable, "-c", CHILD_CODE],
|
||||
cwd=variant.worktree,
|
||||
input=json.dumps(payload),
|
||||
text=True,
|
||||
capture_output=True,
|
||||
check=False,
|
||||
timeout=PROBE_SUBPROCESS_TIMEOUT_SECONDS,
|
||||
)
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
raise RuntimeError(
|
||||
(
|
||||
f"{variant.label} probe timed out for {provider.label} "
|
||||
f"({scenario['name']}) after {PROBE_SUBPROCESS_TIMEOUT_SECONDS} seconds.\n"
|
||||
f"stdout:\n{exc.stdout or ''}\n"
|
||||
f"stderr:\n{exc.stderr or ''}"
|
||||
)
|
||||
) from exc
|
||||
if process.returncode != 0:
|
||||
raise RuntimeError(
|
||||
(
|
||||
|
|
@ -369,7 +466,17 @@ def run_scenario_probe(
|
|||
f"stderr:\n{process.stderr}"
|
||||
)
|
||||
)
|
||||
return json.loads(process.stdout)
|
||||
try:
|
||||
return json.loads(process.stdout)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise RuntimeError(
|
||||
(
|
||||
f"{variant.label} probe produced invalid JSON for {provider.label} "
|
||||
f"({scenario['name']}).\n"
|
||||
f"stdout:\n{process.stdout}\n"
|
||||
f"stderr:\n{process.stderr}"
|
||||
)
|
||||
) from exc
|
||||
|
||||
|
||||
def run_probe(
|
||||
|
|
@ -379,7 +486,18 @@ def run_probe(
|
|||
max_tokens: int,
|
||||
invocation_salt: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Run all requested scenarios for one worktree/provider combination."""
|
||||
"""Run all scenarios for one worktree/provider pair.
|
||||
|
||||
Args:
|
||||
variant: Worktree variant under test.
|
||||
provider: Provider/model pair being evaluated.
|
||||
scenarios: Scenario definitions to execute.
|
||||
max_tokens: Maximum output tokens to request.
|
||||
invocation_salt: Per-process salt generated once at startup.
|
||||
|
||||
Returns:
|
||||
Aggregated results for the variant/provider pair.
|
||||
"""
|
||||
return {
|
||||
"scenarios": [
|
||||
run_scenario_probe(
|
||||
|
|
@ -395,14 +513,28 @@ def run_probe(
|
|||
|
||||
|
||||
def format_metric(value: Any) -> str:
|
||||
"""Render numeric metrics in a stable human-readable form."""
|
||||
"""Render a metric in stable human-readable form.
|
||||
|
||||
Args:
|
||||
value: Metric value to render.
|
||||
|
||||
Returns:
|
||||
String representation of the metric.
|
||||
"""
|
||||
if isinstance(value, float):
|
||||
return f"{value:.2f}"
|
||||
return str(value)
|
||||
|
||||
|
||||
def cache_ratio_pct(call: dict[str, Any]) -> float:
|
||||
"""Compute the percent of input tokens served from cache for a call."""
|
||||
"""Compute the cached-input percentage for one call.
|
||||
|
||||
Args:
|
||||
call: Probe result for one model call.
|
||||
|
||||
Returns:
|
||||
Percentage of input tokens reported as cache reads.
|
||||
"""
|
||||
input_tokens = call["input_tokens"] or 0
|
||||
if input_tokens <= 0:
|
||||
return 0.0
|
||||
|
|
@ -410,7 +542,12 @@ def cache_ratio_pct(call: dict[str, Any]) -> float:
|
|||
|
||||
|
||||
def print_variant_result(variant: VariantSpec, result: dict[str, Any]) -> None:
|
||||
"""Print per-call cache metrics for one variant."""
|
||||
"""Print per-call cache metrics for one variant.
|
||||
|
||||
Args:
|
||||
variant: Worktree variant being reported.
|
||||
result: Probe results for the variant.
|
||||
"""
|
||||
print(f" {variant.label}")
|
||||
for scenario in result["scenarios"]:
|
||||
print(f" [{scenario['name']}]")
|
||||
|
|
@ -431,7 +568,12 @@ def print_delta_summary(
|
|||
baseline: dict[str, Any],
|
||||
candidate: dict[str, Any],
|
||||
) -> None:
|
||||
"""Print candidate-versus-baseline cache and latency deltas."""
|
||||
"""Print candidate-versus-baseline cache and latency deltas.
|
||||
|
||||
Args:
|
||||
baseline: Probe results for the baseline worktree.
|
||||
candidate: Probe results for the candidate worktree.
|
||||
"""
|
||||
baseline_by_name = {scenario["name"]: scenario for scenario in baseline["scenarios"]}
|
||||
candidate_by_name = {scenario["name"]: scenario for scenario in candidate["scenarios"]}
|
||||
print(" delta summary (candidate - baseline)")
|
||||
|
|
@ -463,7 +605,12 @@ def print_delta_summary(
|
|||
|
||||
|
||||
def main() -> None:
|
||||
"""Run the cache comparison probe and optionally write JSON output."""
|
||||
"""Run the probe and optionally write JSON output.
|
||||
|
||||
Raises:
|
||||
ValueError: If one of the provider specifications is malformed.
|
||||
RuntimeError: If any scenario probe fails, times out, or emits invalid JSON.
|
||||
"""
|
||||
args = parse_args()
|
||||
invocation_salt = uuid.uuid4().hex
|
||||
providers = [parse_provider_spec(raw) for raw in args.provider]
|
||||
|
|
|
|||
|
|
@ -62,10 +62,7 @@ def minimal_deriver_prompt(
|
|||
f"""
|
||||
{minimal_deriver_system_prompt(peer_id)}
|
||||
|
||||
Messages to analyze:
|
||||
<messages>
|
||||
{messages}
|
||||
</messages>
|
||||
{minimal_deriver_user_prompt(messages)}
|
||||
"""
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -124,6 +124,48 @@ def _normalize_cacheable_system_content(content: Any) -> list[dict[str, Any]]:
|
|||
return []
|
||||
|
||||
|
||||
def _build_gemini_contents_from_messages(
|
||||
messages: list[dict[str, Any]],
|
||||
) -> tuple[ContentListUnionDict, str | None]:
|
||||
"""Convert generic chat messages into Gemini contents and system instruction."""
|
||||
system_instruction_parts: list[str] = []
|
||||
gemini_contents: list[dict[str, Any]] = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
|
||||
if role == "system":
|
||||
if isinstance(msg.get("content"), str):
|
||||
system_instruction_parts.append(msg["content"])
|
||||
continue
|
||||
|
||||
if role == "assistant":
|
||||
role = "model"
|
||||
|
||||
if isinstance(msg.get("content"), str):
|
||||
gemini_contents.append({"role": role, "parts": [{"text": msg["content"]}]})
|
||||
continue
|
||||
|
||||
if isinstance(msg.get("parts"), list):
|
||||
msg_copy: dict[str, Any] = dict(msg)
|
||||
msg_copy["role"] = role
|
||||
gemini_contents.append(msg_copy)
|
||||
continue
|
||||
|
||||
if isinstance(msg.get("content"), list):
|
||||
text_parts: list[dict[str, str]] = []
|
||||
for block in cast(list[dict[str, Any]], msg["content"]):
|
||||
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
||||
text_parts.append({"text": block["text"]})
|
||||
if text_parts:
|
||||
gemini_contents.append({"role": role, "parts": text_parts})
|
||||
|
||||
system_instruction = (
|
||||
"\n\n".join(system_instruction_parts) if system_instruction_parts else None
|
||||
)
|
||||
return cast(ContentListUnionDict, gemini_contents), system_instruction
|
||||
|
||||
|
||||
def _is_tool_use_message(msg: dict[str, Any]) -> bool:
|
||||
"""Check if a message contains tool calls (any format)."""
|
||||
# Anthropic format: content is a list with tool_use blocks
|
||||
|
|
@ -2158,52 +2200,11 @@ async def honcho_llm_call_inner(
|
|||
|
||||
# Use messages if provided, otherwise use prompt
|
||||
if messages:
|
||||
# Extract system messages for system_instruction parameter
|
||||
# Gemini doesn't support system role in contents - it causes
|
||||
# consecutive user messages which results in empty responses
|
||||
for msg in messages:
|
||||
if msg.get("role") == "system":
|
||||
if isinstance(msg.get("content"), str):
|
||||
system_messages.append(msg["content"])
|
||||
else:
|
||||
non_system_messages.append(msg)
|
||||
|
||||
# Add system instruction if present
|
||||
if system_messages:
|
||||
gemini_config["system_instruction"] = "\n\n".join(
|
||||
system_messages
|
||||
)
|
||||
|
||||
# Convert non-system messages to Google format
|
||||
gemini_contents: list[dict[str, Any]] = []
|
||||
for msg in non_system_messages:
|
||||
# Map roles to Google's expected values (user, model)
|
||||
role = msg.get("role", "user")
|
||||
if role == "assistant":
|
||||
role = "model"
|
||||
|
||||
# Handle different content formats
|
||||
if isinstance(msg.get("content"), str):
|
||||
# Simple string content
|
||||
gemini_contents.append(
|
||||
{"role": role, "parts": [{"text": msg["content"]}]}
|
||||
)
|
||||
elif isinstance(msg.get("parts"), list):
|
||||
# Already in Google format (from tool calling loop)
|
||||
# But still need to ensure role is correct
|
||||
msg_copy = msg.copy()
|
||||
msg_copy["role"] = role
|
||||
gemini_contents.append(msg_copy)
|
||||
elif isinstance(msg.get("content"), list):
|
||||
# Content is a list of parts (Anthropic format) - skip for now
|
||||
# This shouldn't happen with Google provider in tool loop
|
||||
continue
|
||||
else:
|
||||
# Empty or unknown format, skip
|
||||
continue
|
||||
contents: ContentListUnionDict = cast(
|
||||
ContentListUnionDict, gemini_contents
|
||||
contents, system_instruction = _build_gemini_contents_from_messages(
|
||||
messages
|
||||
)
|
||||
if system_instruction:
|
||||
gemini_config["system_instruction"] = system_instruction
|
||||
else:
|
||||
contents = prompt
|
||||
|
||||
|
|
@ -2290,9 +2291,18 @@ async def honcho_llm_call_inner(
|
|||
gemini_config["response_mime_type"] = "application/json"
|
||||
gemini_config["response_schema"] = response_model
|
||||
|
||||
if messages:
|
||||
contents, system_instruction = _build_gemini_contents_from_messages(
|
||||
messages
|
||||
)
|
||||
if system_instruction:
|
||||
gemini_config["system_instruction"] = system_instruction
|
||||
else:
|
||||
contents = prompt
|
||||
|
||||
gemini_response = await client.aio.models.generate_content(
|
||||
model=model,
|
||||
contents=prompt,
|
||||
contents=contents,
|
||||
config=cast(GenerateContentConfigDict, gemini_config), # pyright: ignore[reportInvalidCast]
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ from src.utils.clients import (
|
|||
honcho_llm_call,
|
||||
honcho_llm_call_inner,
|
||||
)
|
||||
from src.utils.representation import ExplicitObservationBase, PromptRepresentation
|
||||
|
||||
|
||||
class SampleTestModel(BaseModel):
|
||||
|
|
@ -783,6 +784,58 @@ class TestGoogleClient:
|
|||
assert config["response_schema"] == SampleTestModel
|
||||
assert config["max_output_tokens"] == 100
|
||||
|
||||
async def test_google_response_model_uses_messages_with_honcho_llm_call(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""Structured Gemini calls should honor system+user messages when provided."""
|
||||
from google import genai
|
||||
|
||||
mock_client = Mock(spec=genai.Client)
|
||||
mock_response = Mock()
|
||||
mock_response.parsed = PromptRepresentation(
|
||||
explicit=[ExplicitObservationBase(content="Alice likes tea")]
|
||||
)
|
||||
mock_finish_reason = Mock()
|
||||
mock_finish_reason.name = "STOP"
|
||||
mock_response.candidates = [Mock(finish_reason=mock_finish_reason)]
|
||||
mock_usage_metadata = Mock()
|
||||
mock_usage_metadata.prompt_token_count = 12
|
||||
mock_usage_metadata.candidates_token_count = 8
|
||||
mock_response.usage_metadata = mock_usage_metadata
|
||||
mock_aio = Mock()
|
||||
mock_aio.models.generate_content = AsyncMock(return_value=mock_response)
|
||||
mock_client.aio = mock_aio
|
||||
|
||||
monkeypatch.setattr(settings.DERIVER, "PROVIDER", "google")
|
||||
monkeypatch.setattr(settings.DERIVER, "MODEL", "gemini-1.5-pro")
|
||||
|
||||
with patch.dict(CLIENTS, {"google": mock_client}):
|
||||
response = await honcho_llm_call(
|
||||
llm_settings=settings.DERIVER,
|
||||
prompt="",
|
||||
max_tokens=100,
|
||||
response_model=PromptRepresentation,
|
||||
messages=[
|
||||
{"role": "system", "content": "stable deriver instructions"},
|
||||
{"role": "user", "content": "Messages to analyze:\n<messages>hello</messages>"},
|
||||
],
|
||||
)
|
||||
|
||||
assert isinstance(response, HonchoLLMCallResponse)
|
||||
assert isinstance(response.content, PromptRepresentation)
|
||||
assert response.content.explicit[0].content == "Alice likes tea"
|
||||
|
||||
call_args = mock_aio.models.generate_content.call_args
|
||||
assert call_args is not None
|
||||
assert call_args.kwargs["contents"] == [
|
||||
{"role": "user", "parts": [{"text": "Messages to analyze:\n<messages>hello</messages>"}]}
|
||||
]
|
||||
config = call_args.kwargs["config"]
|
||||
assert config["system_instruction"] == "stable deriver instructions"
|
||||
assert config["response_schema"] == PromptRepresentation
|
||||
assert config["response_mime_type"] == "application/json"
|
||||
|
||||
async def test_google_streaming(self):
|
||||
"""Test Google/Gemini streaming response"""
|
||||
from google import genai
|
||||
|
|
|
|||
Loading…
Reference in New Issue