Optimize prompt prefix caching
This commit is contained in:
parent
24f94f3ff8
commit
1b99469335
|
|
@ -0,0 +1,415 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Compare prompt-prefix cache behavior between two Honcho worktrees.
|
||||
|
||||
This is intentionally a narrow probe. It does not start Honcho servers, touch the
|
||||
database, or exercise Hermes/session state. Instead, it imports each worktree's
|
||||
`src.utils.clients.honcho_llm_call_inner` and runs a few controlled message
|
||||
patterns against live providers.
|
||||
|
||||
The important scenario is `change_history`: the first system block stays stable
|
||||
while the second rolling system block changes. The candidate branch should retain
|
||||
more cache reuse there because it preserves multiple cacheable system blocks
|
||||
instead of flattening them into one blob.
|
||||
|
||||
Example:
|
||||
uv run python scripts/compare_prefix_cache.py \
|
||||
--baseline-worktree /path/to/honcho-main \
|
||||
--candidate-worktree /path/to/honcho-branch \
|
||||
--provider anthropic-haiku=anthropic:claude-haiku-4-5 \
|
||||
--provider openrouter-haiku=custom:anthropic/claude-haiku-4.5 \
|
||||
--provider openai-mini=openai:gpt-4.1-mini
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
CHILD_CODE = r"""
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
|
||||
payload = json.loads(sys.argv[1])
|
||||
|
||||
from src.utils.clients import honcho_llm_call_inner
|
||||
|
||||
|
||||
async def run() -> None:
|
||||
results = []
|
||||
for scenario in payload["scenarios"]:
|
||||
calls = []
|
||||
for call in scenario["calls"]:
|
||||
start = time.perf_counter()
|
||||
response = await honcho_llm_call_inner(
|
||||
provider=payload["provider"],
|
||||
model=payload["model"],
|
||||
prompt="",
|
||||
max_tokens=payload["max_tokens"],
|
||||
temperature=0,
|
||||
messages=call["messages"],
|
||||
)
|
||||
elapsed_ms = (time.perf_counter() - start) * 1000
|
||||
calls.append(
|
||||
{
|
||||
"label": call["label"],
|
||||
"duration_ms": elapsed_ms,
|
||||
"input_tokens": response.input_tokens,
|
||||
"output_tokens": response.output_tokens,
|
||||
"cache_creation_input_tokens": response.cache_creation_input_tokens,
|
||||
"cache_read_input_tokens": response.cache_read_input_tokens,
|
||||
"finish_reasons": response.finish_reasons,
|
||||
"content_preview": (response.content or "")[:120],
|
||||
}
|
||||
)
|
||||
results.append({"name": scenario["name"], "calls": calls})
|
||||
|
||||
print(json.dumps({"scenarios": results}))
|
||||
|
||||
|
||||
asyncio.run(run())
|
||||
"""
|
||||
|
||||
|
||||
BASE_PREFIX = "\n".join(
|
||||
[
|
||||
"You are Honcho's memory-backed reasoning layer.",
|
||||
"Answer precisely, prefer explicit dates, and preserve user-specific facts.",
|
||||
"Treat the following policy statements as durable background instructions.",
|
||||
]
|
||||
+ [
|
||||
f"Policy {i}: Keep stable user preferences and constraints explicit in memory-aware answers."
|
||||
for i in range(1, 121)
|
||||
]
|
||||
)
|
||||
|
||||
BASE_PREFIX_VARIANT = "\n".join(
|
||||
[
|
||||
"You are Honcho's memory-backed reasoning layer.",
|
||||
"Answer precisely, prefer explicit dates, and preserve user-specific facts.",
|
||||
"Treat the following policy statements as durable background instructions.",
|
||||
]
|
||||
+ [
|
||||
f"Policy {i}: Emphasize durable preferences, deadlines, and factual constraints in every answer."
|
||||
for i in range(1, 121)
|
||||
]
|
||||
)
|
||||
|
||||
ROLLING_HISTORY_A = "\n".join(
|
||||
[
|
||||
"Session history snapshot A:",
|
||||
"The user usually drinks green tea on weekdays and espresso on Sundays.",
|
||||
"The user moved a product launch deadline from April 25, 2026 to April 22, 2026.",
|
||||
"The user prefers short bullet points and exact dates for updates.",
|
||||
]
|
||||
+ [
|
||||
f"History line {i}: The user mentioned project detail {i} while discussing the hermes-memory rollout."
|
||||
for i in range(1, 121)
|
||||
]
|
||||
)
|
||||
|
||||
ROLLING_HISTORY_B = "\n".join(
|
||||
[
|
||||
"Session history snapshot B:",
|
||||
"The user usually drinks green tea on weekdays and espresso on Sundays.",
|
||||
"The user moved a product launch deadline from April 25, 2026 to April 22, 2026.",
|
||||
"The user prefers short bullet points and exact dates for updates.",
|
||||
]
|
||||
+ [
|
||||
f"History line {i}: The user mentioned project detail {i} while discussing the prefix-cache rollout."
|
||||
for i in range(1, 121)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProviderSpec:
|
||||
label: str
|
||||
provider: str
|
||||
model: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class VariantSpec:
|
||||
label: str
|
||||
worktree: Path
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Compare prompt-prefix cache behavior between two Honcho worktrees."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--baseline-worktree",
|
||||
required=True,
|
||||
type=Path,
|
||||
help="Path to the baseline Honcho worktree, typically main.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--candidate-worktree",
|
||||
required=True,
|
||||
type=Path,
|
||||
help="Path to the candidate Honcho worktree.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--provider",
|
||||
action="append",
|
||||
required=True,
|
||||
help=(
|
||||
"Provider spec in the form label=provider:model. "
|
||||
"Example: anthropic-haiku=anthropic:claude-haiku-4-5"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-tokens",
|
||||
type=int,
|
||||
default=128,
|
||||
help="Max output tokens for each probe call.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--scenario",
|
||||
action="append",
|
||||
choices=["repeat_exact", "change_user", "change_history", "change_base"],
|
||||
help="Optional scenario filter. Defaults to all scenarios.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-json",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Optional path to write the raw comparison output as JSON.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def parse_provider_spec(raw: str) -> ProviderSpec:
|
||||
if "=" not in raw or ":" not in raw:
|
||||
raise ValueError(
|
||||
f"Invalid provider spec {raw!r}. Expected label=provider:model"
|
||||
)
|
||||
label, provider_model = raw.split("=", 1)
|
||||
provider, model = provider_model.split(":", 1)
|
||||
return ProviderSpec(label=label, provider=provider, model=model)
|
||||
|
||||
|
||||
def build_messages(base_prefix: str, rolling_history: str, user_query: str) -> list[dict[str, str]]:
|
||||
return [
|
||||
{"role": "system", "content": base_prefix},
|
||||
{"role": "system", "content": rolling_history},
|
||||
{"role": "user", "content": user_query},
|
||||
]
|
||||
|
||||
|
||||
def build_scenarios(selected: set[str] | None) -> list[dict[str, Any]]:
|
||||
scenario_defs = [
|
||||
{
|
||||
"name": "repeat_exact",
|
||||
"calls": [
|
||||
{
|
||||
"label": "cold",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_A,
|
||||
"What is the user's preferred morning drink schedule?",
|
||||
),
|
||||
},
|
||||
{
|
||||
"label": "warm_same",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_A,
|
||||
"What is the user's preferred morning drink schedule?",
|
||||
),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"name": "change_user",
|
||||
"calls": [
|
||||
{
|
||||
"label": "cold",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_A,
|
||||
"What is the user's preferred morning drink schedule?",
|
||||
),
|
||||
},
|
||||
{
|
||||
"label": "warm_user_changed",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_A,
|
||||
"What exact launch date should be remembered for the user?",
|
||||
),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"name": "change_history",
|
||||
"calls": [
|
||||
{
|
||||
"label": "cold",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_A,
|
||||
"Summarize the user's communication preference in one sentence.",
|
||||
),
|
||||
},
|
||||
{
|
||||
"label": "warm_history_changed",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_B,
|
||||
"Summarize the user's communication preference in one sentence.",
|
||||
),
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"name": "change_base",
|
||||
"calls": [
|
||||
{
|
||||
"label": "cold",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX,
|
||||
ROLLING_HISTORY_A,
|
||||
"What city is the user considering for a move?",
|
||||
),
|
||||
},
|
||||
{
|
||||
"label": "warm_base_changed",
|
||||
"messages": build_messages(
|
||||
BASE_PREFIX_VARIANT,
|
||||
ROLLING_HISTORY_A,
|
||||
"What city is the user considering for a move?",
|
||||
),
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
if not selected:
|
||||
return scenario_defs
|
||||
return [scenario for scenario in scenario_defs if scenario["name"] in selected]
|
||||
|
||||
|
||||
def run_probe(
|
||||
variant: VariantSpec,
|
||||
provider: ProviderSpec,
|
||||
scenarios: list[dict[str, Any]],
|
||||
max_tokens: int,
|
||||
) -> dict[str, Any]:
|
||||
payload = {
|
||||
"provider": provider.provider,
|
||||
"model": provider.model,
|
||||
"max_tokens": max_tokens,
|
||||
"scenarios": scenarios,
|
||||
}
|
||||
process = subprocess.run(
|
||||
[sys.executable, "-c", CHILD_CODE, json.dumps(payload)],
|
||||
cwd=variant.worktree,
|
||||
text=True,
|
||||
capture_output=True,
|
||||
check=False,
|
||||
)
|
||||
if process.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"{variant.label} probe failed for {provider.label}.\n"
|
||||
f"stdout:\n{process.stdout}\n"
|
||||
f"stderr:\n{process.stderr}"
|
||||
)
|
||||
return json.loads(process.stdout)
|
||||
|
||||
|
||||
def format_metric(value: Any) -> str:
|
||||
if isinstance(value, float):
|
||||
return f"{value:.2f}"
|
||||
return str(value)
|
||||
|
||||
|
||||
def print_variant_result(variant: VariantSpec, result: dict[str, Any]) -> None:
|
||||
print(f" {variant.label}")
|
||||
for scenario in result["scenarios"]:
|
||||
print(f" [{scenario['name']}]")
|
||||
for call in scenario["calls"]:
|
||||
print(
|
||||
" "
|
||||
f"{call['label']:<18} "
|
||||
f"read={format_metric(call['cache_read_input_tokens']):>8} "
|
||||
f"create={format_metric(call['cache_creation_input_tokens']):>8} "
|
||||
f"input={format_metric(call['input_tokens']):>8} "
|
||||
f"ms={format_metric(call['duration_ms']):>8}"
|
||||
)
|
||||
|
||||
|
||||
def print_delta_summary(
|
||||
baseline: dict[str, Any],
|
||||
candidate: dict[str, Any],
|
||||
) -> None:
|
||||
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)")
|
||||
for name in baseline_by_name:
|
||||
base_calls = baseline_by_name[name]["calls"]
|
||||
cand_calls = candidate_by_name[name]["calls"]
|
||||
if len(base_calls) < 2 or len(cand_calls) < 2:
|
||||
continue
|
||||
base_warm = base_calls[1]
|
||||
cand_warm = cand_calls[1]
|
||||
read_delta = (
|
||||
cand_warm["cache_read_input_tokens"] - base_warm["cache_read_input_tokens"]
|
||||
)
|
||||
create_delta = (
|
||||
cand_warm["cache_creation_input_tokens"]
|
||||
- base_warm["cache_creation_input_tokens"]
|
||||
)
|
||||
latency_delta = cand_warm["duration_ms"] - base_warm["duration_ms"]
|
||||
print(
|
||||
" "
|
||||
f"{name:<16} "
|
||||
f"read_delta={read_delta:+8.2f} "
|
||||
f"create_delta={create_delta:+8.2f} "
|
||||
f"warm_latency_delta_ms={latency_delta:+8.2f}"
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
providers = [parse_provider_spec(raw) for raw in args.provider]
|
||||
baseline = VariantSpec("baseline", args.baseline_worktree.resolve())
|
||||
candidate = VariantSpec("candidate", args.candidate_worktree.resolve())
|
||||
scenarios = build_scenarios(set(args.scenario) if args.scenario else None)
|
||||
|
||||
all_results: dict[str, Any] = {"providers": []}
|
||||
|
||||
for provider in providers:
|
||||
print("=" * 100)
|
||||
print(f"Provider {provider.label}: {provider.provider}/{provider.model}")
|
||||
print("=" * 100)
|
||||
baseline_result = run_probe(baseline, provider, scenarios, args.max_tokens)
|
||||
candidate_result = run_probe(candidate, provider, scenarios, args.max_tokens)
|
||||
print_variant_result(baseline, baseline_result)
|
||||
print_variant_result(candidate, candidate_result)
|
||||
print_delta_summary(baseline_result, candidate_result)
|
||||
all_results["providers"].append(
|
||||
{
|
||||
"label": provider.label,
|
||||
"provider": provider.provider,
|
||||
"model": provider.model,
|
||||
"baseline": baseline_result,
|
||||
"candidate": candidate_result,
|
||||
}
|
||||
)
|
||||
print()
|
||||
|
||||
if args.output_json:
|
||||
args.output_json.write_text(json.dumps(all_results, indent=2) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -22,7 +22,11 @@ from src.utils.formatting import format_new_turn_with_timestamp
|
|||
from src.utils.representation import PromptRepresentation, Representation
|
||||
from src.utils.tokens import track_deriver_input_tokens
|
||||
|
||||
from .prompts import estimate_minimal_deriver_prompt_tokens, minimal_deriver_prompt
|
||||
from .prompts import (
|
||||
estimate_minimal_deriver_prompt_tokens,
|
||||
minimal_deriver_system_prompt,
|
||||
minimal_deriver_user_prompt,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -107,9 +111,6 @@ async def process_representation_tasks_batch(
|
|||
},
|
||||
)
|
||||
|
||||
# Build prompt
|
||||
prompt = minimal_deriver_prompt(peer_id=observed, messages=formatted_messages)
|
||||
|
||||
context_prep_duration = (time.perf_counter() - overall_start) * 1000
|
||||
accumulate_metric(
|
||||
f"minimal_deriver_{latest_message.id}_{observed}",
|
||||
|
|
@ -125,7 +126,7 @@ async def process_representation_tasks_batch(
|
|||
llm_start = time.perf_counter()
|
||||
response = await honcho_llm_call(
|
||||
llm_settings=settings.DERIVER,
|
||||
prompt=prompt,
|
||||
prompt="",
|
||||
max_tokens=max_tokens,
|
||||
track_name="Minimal Deriver",
|
||||
response_model=PromptRepresentation,
|
||||
|
|
@ -138,6 +139,13 @@ async def process_representation_tasks_batch(
|
|||
enable_retry=True,
|
||||
retry_attempts=3,
|
||||
trace_name="minimal_deriver",
|
||||
messages=[
|
||||
{"role": "system", "content": minimal_deriver_system_prompt(observed)},
|
||||
{
|
||||
"role": "user",
|
||||
"content": minimal_deriver_user_prompt(formatted_messages),
|
||||
},
|
||||
],
|
||||
)
|
||||
llm_duration = (time.perf_counter() - llm_start) * 1000
|
||||
|
||||
|
|
|
|||
|
|
@ -11,20 +11,8 @@ from inspect import cleandoc as c
|
|||
from src.utils.tokens import estimate_tokens
|
||||
|
||||
|
||||
def minimal_deriver_prompt(
|
||||
peer_id: str,
|
||||
messages: str,
|
||||
) -> str:
|
||||
"""
|
||||
Generate minimal prompt for fast observation extraction.
|
||||
|
||||
Args:
|
||||
peer_id: The ID of the user being analyzed.
|
||||
messages: All messages in the range (interleaving messages and new turns combined).
|
||||
|
||||
Returns:
|
||||
Formatted prompt string for observation extraction.
|
||||
"""
|
||||
def minimal_deriver_system_prompt(peer_id: str) -> str:
|
||||
"""Generate the cacheable instructions for observation extraction."""
|
||||
return c(
|
||||
f"""
|
||||
Analyze messages from {peer_id} to extract **explicit atomic facts** about them.
|
||||
|
|
@ -44,6 +32,35 @@ EXAMPLES:
|
|||
- EXPLICIT: "I just had my 25th birthday last Saturday" → "{peer_id} is 25 years old", "{peer_id}'s birthday is June 21st"
|
||||
- EXPLICIT: "I took my dog for a walk in NYC" → "{peer_id} has a dog", "{peer_id} lives in NYC"
|
||||
- EXPLICIT: "{peer_id} attended college" + general knowledge → "{peer_id} completed high school or equivalent"
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def minimal_deriver_user_prompt(messages: str) -> str:
|
||||
"""Generate the per-request message payload for observation extraction."""
|
||||
return c(
|
||||
f"""
|
||||
Messages to analyze:
|
||||
<messages>
|
||||
{messages}
|
||||
</messages>
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def minimal_deriver_prompt(
|
||||
peer_id: str,
|
||||
messages: str,
|
||||
) -> str:
|
||||
"""
|
||||
Generate the combined prompt for fast observation extraction.
|
||||
|
||||
Prefer `minimal_deriver_system_prompt()` plus `minimal_deriver_user_prompt()`
|
||||
when making LLM calls so the instructions can be cached independently.
|
||||
"""
|
||||
return c(
|
||||
f"""
|
||||
{minimal_deriver_system_prompt(peer_id)}
|
||||
|
||||
Messages to analyze:
|
||||
<messages>
|
||||
|
|
@ -57,9 +74,11 @@ Messages to analyze:
|
|||
def estimate_minimal_deriver_prompt_tokens() -> int:
|
||||
"""Estimate base prompt tokens (cached)."""
|
||||
try:
|
||||
prompt = minimal_deriver_prompt(
|
||||
peer_id="",
|
||||
messages="",
|
||||
prompt = "\n\n".join(
|
||||
[
|
||||
minimal_deriver_system_prompt(peer_id=""),
|
||||
minimal_deriver_user_prompt(messages=""),
|
||||
]
|
||||
)
|
||||
return estimate_tokens(prompt)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -141,8 +141,9 @@ class DialecticAgent:
|
|||
"</session_history>"
|
||||
)
|
||||
|
||||
# Append session history to the system prompt
|
||||
self.messages[0]["content"] += session_history_section
|
||||
# Keep session history in its own system message so the stable base
|
||||
# instructions can be cached independently of rolling session context.
|
||||
self.messages.append({"role": "system", "content": session_history_section})
|
||||
|
||||
async def _prefetch_relevant_observations(self, query: str) -> str | None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -92,6 +92,37 @@ def count_message_tokens(messages: list[dict[str, Any]]) -> int:
|
|||
return total
|
||||
|
||||
|
||||
def _cacheable_text_block(text: str) -> dict[str, Any]:
|
||||
"""Create a text content block that participates in prompt caching."""
|
||||
return {
|
||||
"type": "text",
|
||||
"text": text,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
|
||||
|
||||
def _normalize_cacheable_system_content(content: Any) -> list[dict[str, Any]]:
|
||||
"""Normalize system content into cacheable text blocks when possible."""
|
||||
if isinstance(content, str):
|
||||
return [_cacheable_text_block(content)]
|
||||
|
||||
if isinstance(content, list):
|
||||
normalized_blocks: list[dict[str, Any]] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
|
||||
normalized_block = dict(block)
|
||||
if normalized_block.get("type") == "text":
|
||||
normalized_block.setdefault(
|
||||
"cache_control", {"type": "ephemeral"}
|
||||
)
|
||||
normalized_blocks.append(normalized_block)
|
||||
return normalized_blocks
|
||||
|
||||
return []
|
||||
|
||||
|
||||
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
|
||||
|
|
@ -1678,7 +1709,7 @@ async def honcho_llm_call_inner(
|
|||
# Remove stream parameter for non-streaming calls as some providers don't accept it
|
||||
params.pop("stream", None)
|
||||
|
||||
system_messages: list[str] = []
|
||||
system_messages: list[Any] = []
|
||||
non_system_messages: list[dict[str, Any]] = []
|
||||
|
||||
match client:
|
||||
|
|
@ -1703,13 +1734,14 @@ async def honcho_llm_call_inner(
|
|||
# Add system parameter if there are system messages
|
||||
# Use cache_control for prompt caching
|
||||
if system_messages:
|
||||
anthropic_params["system"] = [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "\n\n".join(system_messages),
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
]
|
||||
anthropic_system_blocks: list[dict[str, Any]] = []
|
||||
for system_message in system_messages:
|
||||
anthropic_system_blocks.extend(
|
||||
_normalize_cacheable_system_content(system_message)
|
||||
)
|
||||
|
||||
if anthropic_system_blocks:
|
||||
anthropic_params["system"] = anthropic_system_blocks
|
||||
|
||||
# Add tools if provided
|
||||
if tools:
|
||||
|
|
@ -1853,22 +1885,19 @@ async def honcho_llm_call_inner(
|
|||
if provider == "custom":
|
||||
processed_messages = []
|
||||
for msg in params["messages"]:
|
||||
if msg.get("role") == "system" and isinstance(
|
||||
msg.get("content"), str
|
||||
):
|
||||
# Convert system message to content block format with cache_control
|
||||
processed_messages.append(
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": msg["content"],
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}
|
||||
if msg.get("role") == "system":
|
||||
cacheable_content = _normalize_cacheable_system_content(
|
||||
msg.get("content")
|
||||
)
|
||||
if cacheable_content:
|
||||
processed_messages.append(
|
||||
{
|
||||
**msg,
|
||||
"content": cacheable_content,
|
||||
}
|
||||
)
|
||||
else:
|
||||
processed_messages.append(msg)
|
||||
else:
|
||||
processed_messages.append(msg)
|
||||
|
||||
|
|
@ -2373,22 +2402,19 @@ async def handle_streaming_response(
|
|||
case AsyncAnthropic():
|
||||
# Anthropic requires system messages as a top-level parameter
|
||||
messages = params["messages"]
|
||||
system_content = "\n\n".join(
|
||||
m["content"] for m in messages if m.get("role") == "system"
|
||||
)
|
||||
system_blocks: list[dict[str, Any]] = []
|
||||
for message in messages:
|
||||
if message.get("role") == "system":
|
||||
system_blocks.extend(
|
||||
_normalize_cacheable_system_content(message.get("content"))
|
||||
)
|
||||
anthropic_params: dict[str, Any] = {
|
||||
"model": params["model"],
|
||||
"max_tokens": params["max_tokens"],
|
||||
"messages": [m for m in messages if m.get("role") != "system"],
|
||||
}
|
||||
if system_content:
|
||||
anthropic_params["system"] = [
|
||||
{
|
||||
"type": "text",
|
||||
"text": system_content,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
]
|
||||
if system_blocks:
|
||||
anthropic_params["system"] = system_blocks
|
||||
|
||||
# For response models, we need to request JSON and parse manually
|
||||
# Note: Streaming with response_model is not ideal but we'll accumulate and parse at the end
|
||||
|
|
|
|||
|
|
@ -91,13 +91,9 @@ class SummaryType(Enum):
|
|||
LONG = "honcho_chat_summary_long"
|
||||
|
||||
|
||||
def short_summary_prompt(
|
||||
formatted_messages: str,
|
||||
output_words: int,
|
||||
previous_summary_text: str,
|
||||
) -> str:
|
||||
"""Generate the short summary prompt."""
|
||||
return c(f"""
|
||||
def short_summary_system_prompt() -> str:
|
||||
"""Generate cacheable instructions for short summaries."""
|
||||
return c("""
|
||||
You are a system that summarizes parts of a conversation to create a concise and accurate summary. Focus on capturing:
|
||||
|
||||
1. Key facts and information shared (**Capture as many explicit facts as possible**)
|
||||
|
|
@ -110,7 +106,71 @@ If there is a previous summary, ALWAYS make your new summary inclusive of both i
|
|||
Provide a concise, factual summary that captures the essence of the conversation. Your summary should be detailed enough to serve as context for future messages, but brief enough to be helpful. Prefer a thorough chronological narrative over a list of bullet points.
|
||||
|
||||
Return only the summary without any explanation or meta-commentary.
|
||||
""")
|
||||
|
||||
|
||||
def short_summary_user_prompt(
|
||||
formatted_messages: str,
|
||||
output_words: int,
|
||||
previous_summary_text: str,
|
||||
) -> str:
|
||||
"""Generate the per-request payload for short summaries."""
|
||||
return c(f"""
|
||||
<previous_summary>
|
||||
{previous_summary_text}
|
||||
</previous_summary>
|
||||
|
||||
<conversation>
|
||||
{formatted_messages}
|
||||
</conversation>
|
||||
|
||||
Hard limit: {output_words} words maximum. If needed, drop lower-priority detail to stay within the limit.
|
||||
""")
|
||||
|
||||
|
||||
def short_summary_prompt(
|
||||
formatted_messages: str,
|
||||
output_words: int,
|
||||
previous_summary_text: str,
|
||||
) -> str:
|
||||
"""Generate the combined short summary prompt."""
|
||||
return "\n\n".join(
|
||||
[
|
||||
short_summary_system_prompt(),
|
||||
short_summary_user_prompt(
|
||||
formatted_messages, output_words, previous_summary_text
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def long_summary_system_prompt() -> str:
|
||||
"""Generate cacheable instructions for long summaries."""
|
||||
return c("""
|
||||
You are a system that creates thorough, comprehensive summaries of conversations. Focus on capturing:
|
||||
|
||||
1. Key facts and information shared (**Capture as many explicit facts as possible**)
|
||||
2. User preferences, opinions, and questions
|
||||
3. Important context and requests
|
||||
4. Core topics discussed in detail
|
||||
5. User's apparent emotional state and personality traits
|
||||
6. Important themes and patterns across the conversation
|
||||
|
||||
If there is a previous summary, ALWAYS make your new summary inclusive of both it and the new messages, therefore capturing the ENTIRE conversation. Prioritize key facts across the entire conversation.
|
||||
|
||||
Provide a thorough and detailed summary that captures the essence of the conversation. Your summary should serve as a comprehensive record of the important information in this conversation. Prefer an exhaustive chronological narrative over a list of bullet points.
|
||||
|
||||
Return only the summary without any explanation or meta-commentary.
|
||||
""")
|
||||
|
||||
|
||||
def long_summary_user_prompt(
|
||||
formatted_messages: str,
|
||||
output_words: int,
|
||||
previous_summary_text: str,
|
||||
) -> str:
|
||||
"""Generate the per-request payload for long summaries."""
|
||||
return c(f"""
|
||||
<previous_summary>
|
||||
{previous_summary_text}
|
||||
</previous_summary>
|
||||
|
|
@ -128,33 +188,15 @@ def long_summary_prompt(
|
|||
output_words: int,
|
||||
previous_summary_text: str,
|
||||
) -> str:
|
||||
"""Generate the long summary prompt."""
|
||||
return c(f"""
|
||||
You are a system that creates thorough, comprehensive summaries of conversations. Focus on capturing:
|
||||
|
||||
1. Key facts and information shared (**Capture as many explicit facts as possible**)
|
||||
2. User preferences, opinions, and questions
|
||||
3. Important context and requests
|
||||
4. Core topics discussed in detail
|
||||
5. User's apparent emotional state and personality traits
|
||||
6. Important themes and patterns across the conversation
|
||||
|
||||
If there is a previous summary, ALWAYS make your new summary inclusive of both it and the new messages, therefore capturing the ENTIRE conversation. Prioritize key facts across the entire conversation.
|
||||
|
||||
Provide a thorough and detailed summary that captures the essence of the conversation. Your summary should serve as a comprehensive record of the important information in this conversation. Prefer an exhaustive chronological narrative over a list of bullet points.
|
||||
|
||||
Return only the summary without any explanation or meta-commentary.
|
||||
|
||||
<previous_summary>
|
||||
{previous_summary_text}
|
||||
</previous_summary>
|
||||
|
||||
<conversation>
|
||||
{formatted_messages}
|
||||
</conversation>
|
||||
|
||||
Hard limit: {output_words} words maximum. If needed, drop lower-priority detail to stay within the limit.
|
||||
""")
|
||||
"""Generate the combined long summary prompt."""
|
||||
return "\n\n".join(
|
||||
[
|
||||
long_summary_system_prompt(),
|
||||
long_summary_user_prompt(
|
||||
formatted_messages, output_words, previous_summary_text
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@cache
|
||||
|
|
@ -162,10 +204,15 @@ def estimate_short_summary_prompt_tokens() -> int:
|
|||
"""Estimate tokens for the short summary prompt (without messages/previous_summary)."""
|
||||
try:
|
||||
return estimate_tokens(
|
||||
short_summary_prompt(
|
||||
formatted_messages="",
|
||||
output_words=0,
|
||||
previous_summary_text="",
|
||||
"\n\n".join(
|
||||
[
|
||||
short_summary_system_prompt(),
|
||||
short_summary_user_prompt(
|
||||
formatted_messages="",
|
||||
output_words=0,
|
||||
previous_summary_text="",
|
||||
),
|
||||
]
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
|
|
@ -178,10 +225,15 @@ def estimate_long_summary_prompt_tokens() -> int:
|
|||
"""Estimate tokens for the long summary prompt (without messages/previous_summary)."""
|
||||
try:
|
||||
return estimate_tokens(
|
||||
long_summary_prompt(
|
||||
formatted_messages="",
|
||||
output_words=0,
|
||||
previous_summary_text="",
|
||||
"\n\n".join(
|
||||
[
|
||||
long_summary_system_prompt(),
|
||||
long_summary_user_prompt(
|
||||
formatted_messages="",
|
||||
output_words=0,
|
||||
previous_summary_text="",
|
||||
),
|
||||
]
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
|
|
@ -207,14 +259,19 @@ async def create_short_summary(
|
|||
else:
|
||||
previous_summary_text = "There is no previous summary -- the messages are the beginning of the conversation."
|
||||
|
||||
prompt = short_summary_prompt(
|
||||
formatted_messages, output_words, previous_summary_text
|
||||
)
|
||||
|
||||
return await honcho_llm_call(
|
||||
llm_settings=settings.SUMMARY,
|
||||
prompt=prompt,
|
||||
prompt="",
|
||||
max_tokens=settings.SUMMARY.MAX_TOKENS_SHORT,
|
||||
messages=[
|
||||
{"role": "system", "content": short_summary_system_prompt()},
|
||||
{
|
||||
"role": "user",
|
||||
"content": short_summary_user_prompt(
|
||||
formatted_messages, output_words, previous_summary_text
|
||||
),
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -232,14 +289,19 @@ async def create_long_summary(
|
|||
else:
|
||||
previous_summary_text = "There is no previous summary -- the messages are the beginning of the conversation."
|
||||
|
||||
prompt = long_summary_prompt(
|
||||
formatted_messages, output_words, previous_summary_text
|
||||
)
|
||||
|
||||
return await honcho_llm_call(
|
||||
llm_settings=settings.SUMMARY,
|
||||
prompt=prompt,
|
||||
prompt="",
|
||||
max_tokens=settings.SUMMARY.MAX_TOKENS_LONG,
|
||||
messages=[
|
||||
{"role": "system", "content": long_summary_system_prompt()},
|
||||
{
|
||||
"role": "user",
|
||||
"content": long_summary_user_prompt(
|
||||
formatted_messages, output_words, previous_summary_text
|
||||
),
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -142,6 +142,37 @@ class TestAnthropicClient:
|
|||
assert response.content == "First block\nSecond block"
|
||||
assert response.output_tokens == 8
|
||||
|
||||
async def test_anthropic_preserves_multiple_cacheable_system_blocks(self):
|
||||
"""Anthropic requests should keep separate system blocks cacheable."""
|
||||
|
||||
mock_client = AsyncMock(spec=AsyncAnthropic)
|
||||
mock_response = Mock()
|
||||
mock_response.content = [TextBlock(text="Hello from Anthropic", type="text")]
|
||||
mock_response.usage = Usage(input_tokens=10, output_tokens=5)
|
||||
mock_response.stop_reason = "stop"
|
||||
mock_client.messages.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch.dict(CLIENTS, {"anthropic": mock_client}):
|
||||
await honcho_llm_call_inner(
|
||||
provider="anthropic",
|
||||
model="claude-3-sonnet",
|
||||
prompt="ignored",
|
||||
max_tokens=100,
|
||||
messages=[
|
||||
{"role": "system", "content": "stable instructions"},
|
||||
{"role": "system", "content": "rolling session context"},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
)
|
||||
|
||||
system_blocks = mock_client.messages.create.call_args.kwargs["system"]
|
||||
assert len(system_blocks) == 2
|
||||
assert system_blocks[0]["text"] == "stable instructions"
|
||||
assert system_blocks[1]["text"] == "rolling session context"
|
||||
assert all(
|
||||
block["cache_control"] == {"type": "ephemeral"} for block in system_blocks
|
||||
)
|
||||
|
||||
async def test_anthropic_json_mode(self):
|
||||
"""Test Anthropic with JSON mode"""
|
||||
|
||||
|
|
@ -518,7 +549,52 @@ class TestOpenAIClient:
|
|||
assert chunks[1].content == " world"
|
||||
assert chunks[2].content == ""
|
||||
assert chunks[2].is_done is True
|
||||
assert chunks[2].finish_reasons == ["stop"]
|
||||
|
||||
async def test_custom_provider_system_messages_gain_cache_control(self):
|
||||
"""Custom OpenAI-compatible providers should mark each system message cacheable."""
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
mock_client = AsyncMock(spec=AsyncOpenAI)
|
||||
mock_response = ChatCompletion(
|
||||
id="test-id",
|
||||
object="chat.completion",
|
||||
created=1234567890,
|
||||
model="anthropic/claude-sonnet",
|
||||
choices=[
|
||||
Choice(
|
||||
index=0,
|
||||
message=ChatCompletionMessage(
|
||||
role="assistant", content="Hello from OpenAI"
|
||||
),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
usage=CompletionUsage(
|
||||
prompt_tokens=10, completion_tokens=5, total_tokens=15
|
||||
),
|
||||
)
|
||||
mock_client.chat.completions.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch.dict(CLIENTS, {"custom": mock_client}):
|
||||
await honcho_llm_call_inner(
|
||||
provider="custom",
|
||||
model="anthropic/claude-sonnet",
|
||||
prompt="ignored",
|
||||
max_tokens=100,
|
||||
messages=[
|
||||
{"role": "system", "content": "stable instructions"},
|
||||
{"role": "system", "content": "session history"},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
)
|
||||
|
||||
messages = mock_client.chat.completions.create.call_args.kwargs["messages"]
|
||||
assert messages[0]["role"] == "system"
|
||||
assert messages[1]["role"] == "system"
|
||||
assert messages[0]["content"][0]["text"] == "stable instructions"
|
||||
assert messages[1]["content"][0]["text"] == "session history"
|
||||
assert messages[0]["content"][0]["cache_control"] == {"type": "ephemeral"}
|
||||
assert messages[1]["content"][0]["cache_control"] == {"type": "ephemeral"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -15,6 +15,8 @@ from src.utils.summarizer import (
|
|||
Summary,
|
||||
SummaryType,
|
||||
_create_summary, # pyright: ignore[reportPrivateUsage]
|
||||
create_long_summary,
|
||||
create_short_summary,
|
||||
)
|
||||
|
||||
# Common test arguments for _create_summary
|
||||
|
|
@ -217,3 +219,71 @@ class TestCreateSummary:
|
|||
assert is_fallback is True
|
||||
assert summary["content"] == ""
|
||||
assert summary["token_count"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestSummaryPromptCaching:
|
||||
"""Tests for cache-friendly summary prompt construction."""
|
||||
|
||||
async def test_create_short_summary_uses_system_and_user_messages(self):
|
||||
"""Short summaries should send stable instructions as a system message."""
|
||||
mock_response = HonchoLLMCallResponse(
|
||||
content="summary",
|
||||
input_tokens=10,
|
||||
output_tokens=5,
|
||||
finish_reasons=["STOP"],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"src.utils.summarizer.honcho_llm_call",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_call:
|
||||
response = await create_short_summary(
|
||||
formatted_messages=_FORMATTED_MESSAGES,
|
||||
input_tokens=_INPUT_TOKENS,
|
||||
previous_summary="earlier summary",
|
||||
)
|
||||
|
||||
assert response is mock_response
|
||||
call_kwargs = mock_call.await_args.kwargs
|
||||
assert call_kwargs["prompt"] == ""
|
||||
assert len(call_kwargs["messages"]) == 2
|
||||
assert call_kwargs["messages"][0]["role"] == "system"
|
||||
assert "summarizes parts of a conversation" in call_kwargs["messages"][0][
|
||||
"content"
|
||||
]
|
||||
assert call_kwargs["messages"][1]["role"] == "user"
|
||||
assert "<previous_summary>" in call_kwargs["messages"][1]["content"]
|
||||
assert "<conversation>" in call_kwargs["messages"][1]["content"]
|
||||
|
||||
async def test_create_long_summary_uses_system_and_user_messages(self):
|
||||
"""Long summaries should send stable instructions as a system message."""
|
||||
mock_response = HonchoLLMCallResponse(
|
||||
content="long summary",
|
||||
input_tokens=10,
|
||||
output_tokens=5,
|
||||
finish_reasons=["STOP"],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"src.utils.summarizer.honcho_llm_call",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_call:
|
||||
response = await create_long_summary(
|
||||
formatted_messages=_FORMATTED_MESSAGES,
|
||||
previous_summary="earlier summary",
|
||||
)
|
||||
|
||||
assert response is mock_response
|
||||
call_kwargs = mock_call.await_args.kwargs
|
||||
assert call_kwargs["prompt"] == ""
|
||||
assert len(call_kwargs["messages"]) == 2
|
||||
assert call_kwargs["messages"][0]["role"] == "system"
|
||||
assert "creates thorough, comprehensive summaries" in call_kwargs[
|
||||
"messages"
|
||||
][0]["content"]
|
||||
assert call_kwargs["messages"][1]["role"] == "user"
|
||||
assert "<previous_summary>" in call_kwargs["messages"][1]["content"]
|
||||
assert "<conversation>" in call_kwargs["messages"][1]["content"]
|
||||
|
|
|
|||
Loading…
Reference in New Issue