Merge branch 'simp/r3-27-W4' into simp/integration3

This commit is contained in:
Teknium
2026-09-03 00:11:47 -07:00

View File

@@ -9,6 +9,7 @@ truncation) | full (truncated raw content). See README.md.
from __future__ import annotations
import atexit
import contextlib
import json
import logging
import os
@@ -87,6 +88,15 @@ def _debug(message: str) -> None:
logger.info("Langfuse tracing: %s", message)
@contextlib.contextmanager
def _failsafe(label: str):
"""Swallow + debug-log any exception: telemetry must never block the agent turn."""
try:
yield
except Exception as exc: # pragma: no cover - fail-open
_debug(f"{label} failed: {exc}")
_CAPTURE_MODES = ("metadata", "sanitized", "full")
_DEFAULT_CAPTURE_MODE = "sanitized"
_warned_invalid_capture = False
@@ -98,10 +108,8 @@ def _capture_mode() -> str:
capture more than the operator intended)."""
global _warned_invalid_capture
value = _env("HERMES_LANGFUSE_CAPTURE").lower()
if not value:
return _DEFAULT_CAPTURE_MODE
if value in _CAPTURE_MODES:
return value
if not value or value in _CAPTURE_MODES:
return value or _DEFAULT_CAPTURE_MODE
if not _warned_invalid_capture:
_warned_invalid_capture = True
logger.warning(
@@ -122,45 +130,42 @@ def _redact_secrets(value: str) -> str:
return value
# (types, shape builder) for _describe_content; first match wins (bool handled before).
_CONTENT_SHAPES = (
((int, float), lambda v: {"type": "number"}),
(bytes, lambda v: {"type": "bytes", "length": len(v)}),
(str, lambda v: {"type": "text", "chars": len(v)}),
(dict, lambda v: {"type": "object", "keys": [str(k) for k in list(v.keys())[:20]]}),
((list, tuple, set), lambda v: {"type": "array", "items": len(v)}),
)
def _describe_content(value: Any) -> Any:
"""Metadata-mode stand-in for content: shape and size, never payload."""
if value is None or isinstance(value, bool):
return value
if isinstance(value, (int, float)):
shape = {"type": "number"}
elif isinstance(value, bytes):
shape = {"type": "bytes", "length": len(value)}
elif isinstance(value, str):
shape = {"type": "text", "chars": len(value)}
elif isinstance(value, dict):
shape = {"type": "object", "keys": [str(k) for k in list(value.keys())[:20]]}
elif isinstance(value, (list, tuple, set)):
shape = {"type": "array", "items": len(value)}
else:
shape = {"type": type(value).__name__}
return {"omitted": True, **shape}
shape = next((build(value) for types, build in _CONTENT_SHAPES if isinstance(value, types)), None)
return {"omitted": True, **(shape or {"type": type(value).__name__})}
def _capture_content(value: Any, *, parse_json_strings: bool = False) -> Any:
def _capture_content(value: Any, *, parse_json_strings: bool = False, tool_result_of: Optional[tuple] = None) -> Any:
"""Apply the active capture mode to a CONTENT value.
Only prompt/response text, tool arguments and tool results are content;
metadata fields (provider, model, IDs, counts) stay as-is in every mode.
``tool_result_of=(tool_name, args)`` marks a tool result: JSON strings are
parsed first so a read_file payload can be collapsed to a preview keyed by
the call's ``args``.
"""
if _capture_mode() == "metadata":
return _describe_content(value)
if tool_result_of is not None:
tool_name, args = tool_result_of
value = _maybe_parse_json_string(value) if isinstance(value, str) else value
value, parse_json_strings = _normalize_payload(value, tool_name=tool_name, args=args), True
return _safe_value(value, parse_json_strings=parse_json_strings)
def _capture_tool_result(result: Any, *, tool_name: str, args: Any) -> Any:
"""Capture a tool result: JSON strings are parsed first so a read_file
payload can be collapsed to a preview keyed by the call's ``args``."""
if _capture_mode() == "metadata":
return _describe_content(result)
value = _maybe_parse_json_string(result) if isinstance(result, str) else result
return _safe_value(_normalize_payload(value, tool_name=tool_name, args=args), parse_json_strings=True)
# Sentinel: "_get_langfuse() has tried and failed". Tests reset by reloading
# the module; runtime callers must restart the process after fixing credentials.
_INIT_FAILED = object()
@@ -181,22 +186,19 @@ def _get_langfuse() -> Optional[Langfuse]:
The first build is serialized so racing callers can't each construct a client
and leak the loser's HTTP connection + flush thread."""
global _LANGFUSE_CLIENT
# Fast path — already settled (success or _INIT_FAILED); no lock needed.
if _LANGFUSE_CLIENT is not None:
return None if _LANGFUSE_CLIENT is _INIT_FAILED else _LANGFUSE_CLIENT
with _LANGFUSE_CLIENT_LOCK:
# Re-check: a racing thread may have finished init while we waited.
if _LANGFUSE_CLIENT is not None:
return None if _LANGFUSE_CLIENT is _INIT_FAILED else _LANGFUSE_CLIENT
client = _build_client()
_LANGFUSE_CLIENT = _INIT_FAILED if client is None else client
if client is not None:
# atexit is LIFO: registering AFTER the SDK's constructor means our
# finalizer runs first, so root spans ended there still get flushed
# by the SDK (short-lived processes: kanban workers, chat -q, cron).
atexit.register(_finalize_all_traces)
return client
# Fast path — already settled (success or _INIT_FAILED) needs no lock;
# re-check under it since a racing thread may have finished init.
if _LANGFUSE_CLIENT is None:
with _LANGFUSE_CLIENT_LOCK:
if _LANGFUSE_CLIENT is None:
client = _build_client()
_LANGFUSE_CLIENT = _INIT_FAILED if client is None else client
if client is not None:
# atexit is LIFO: registering AFTER the SDK's constructor means our
# finalizer runs first, so root spans ended there still get flushed
# by the SDK (short-lived processes: kanban workers, chat -q, cron).
atexit.register(_finalize_all_traces)
return None if _LANGFUSE_CLIENT is _INIT_FAILED else _LANGFUSE_CLIENT
def _build_client() -> Optional[Langfuse]:
@@ -209,17 +211,16 @@ def _build_client() -> Optional[Langfuse]:
)
return None
public_key = _env("HERMES_LANGFUSE_PUBLIC_KEY") or _env("LANGFUSE_PUBLIC_KEY")
secret_key = _env("HERMES_LANGFUSE_SECRET_KEY") or _env("LANGFUSE_SECRET_KEY")
public_key, secret_key = (_env(f"HERMES_LANGFUSE_{n}") or _env(f"LANGFUSE_{n}") for n in ("PUBLIC_KEY", "SECRET_KEY"))
if not (public_key and secret_key):
return None
# The SDK does not validate keys at construction; placeholder keys
# would fail silently at flush time (#23823). Warn once here instead.
placeholder_issues = list(filter(None, (
placeholder_issues = [issue for issue in (
_validate_langfuse_key("HERMES_LANGFUSE_PUBLIC_KEY", public_key),
_validate_langfuse_key("HERMES_LANGFUSE_SECRET_KEY", secret_key),
)))
) if issue]
if placeholder_issues:
logger.warning(
"Langfuse plugin: credentials look like placeholders, traces will "
@@ -230,13 +231,10 @@ def _build_client() -> Optional[Langfuse]:
)
return None
kwargs: Dict[str, Any] = {
"public_key": public_key,
"secret_key": secret_key,
"base_url": _env("HERMES_LANGFUSE_BASE_URL") or _env("LANGFUSE_BASE_URL") or "https://cloud.langfuse.com",
}
for key, name in (("environment", "ENV"), ("release", "RELEASE")):
value = _env(f"HERMES_LANGFUSE_{name}") or _env(f"LANGFUSE_{name}")
kwargs: Dict[str, Any] = {"public_key": public_key, "secret_key": secret_key}
for key, name, default in (("base_url", "BASE_URL", "https://cloud.langfuse.com"), ("environment", "ENV", ""),
("release", "RELEASE", "")):
value = _env(f"HERMES_LANGFUSE_{name}") or _env(f"LANGFUSE_{name}") or default
if value:
kwargs[key] = value
sample_rate = _env("HERMES_LANGFUSE_SAMPLE_RATE")
@@ -253,26 +251,17 @@ def _build_client() -> Optional[Langfuse]:
return None
def _scope_prefix(task_id: str, session_id: str) -> str:
if task_id:
return f"task:{task_id}"
if session_id:
return f"session:{session_id}"
return f"thread:{threading.get_ident()}"
def _trace_key(task_id: str, session_id: str, *, turn_id: str = "", api_request_id: str = "") -> str:
"""In-process trace scope key for one agent turn. ``turn_id`` wins over
``api_request_id`` so the turn-level post_llm_call hook (no api_request_id)
resolves to the same key as request-level hooks; a bare ``task_id`` is the
legacy shape from before turn/request scoping."""
scope = f"task:{task_id}" if task_id else f"session:{session_id}" if session_id else f"thread:{threading.get_ident()}"
if turn_id:
return f"{_scope_prefix(task_id, session_id)}:turn:{turn_id}"
return f"{scope}:turn:{turn_id}"
if api_request_id:
return f"{_scope_prefix(task_id, session_id)}:api:{api_request_id}"
if task_id:
return task_id
return _scope_prefix(task_id, session_id)
return f"{scope}:api:{api_request_id}"
return task_id or scope
def _state_for_turn(turn_id: str) -> Optional[TraceState]:
@@ -285,18 +274,14 @@ def _state_for_turn(turn_id: str) -> Optional[TraceState]:
return next((state for key, state in _TRACE_STATE.items() if key.endswith(suffix)), None)
def _redact_data_uri(value: str) -> dict[str, Any]:
header = value.split(",", 1)[0] if "," in value else "data:"
media_type = header[5:].split(";", 1)[0] if header.startswith("data:") else ""
return {"type": "data_uri", "media_type": media_type or None, "omitted": True, "length": len(value)}
def _truncate_text(value: str, max_chars: int) -> Any:
# The SDK decodes data:*;base64 strings as media; a truncated one is
# invalid base64 and logs noisily, so redact the whole URI instead.
prefix = value[:200].lower()
if prefix.startswith("data:") and ";base64," in prefix:
return _redact_data_uri(value)
header = value.split(",", 1)[0] if "," in value else "data:"
media_type = header[5:].split(";", 1)[0] if header.startswith("data:") else ""
return {"type": "data_uri", "media_type": media_type or None, "omitted": True, "length": len(value)}
# Redact BEFORE truncating so a secret straddling the cut cannot leak.
if _capture_mode() == "sanitized":
value = _redact_secrets(value)
@@ -325,23 +310,25 @@ def _maybe_parse_json_string(value: str) -> Any:
return {"data": parsed, hint_key: trailing}
def _parse_read_file_lines(content: str) -> list[dict[str, Any]]:
if not isinstance(content, str) or not content:
return []
matches = [_READ_FILE_LINE_RE.match(raw_line) for raw_line in content.splitlines()]
if not all(matches):
return []
return [{"line": int(m.group(1)), "text": m.group(2)} for m in matches]
def _normalize_read_file_payload(value: dict[str, Any], *, args: Any = None) -> dict[str, Any]:
def _normalize_payload(value: Any, *, tool_name: str = "", args: Any = None) -> Any:
"""Collapse a read_file result (line-numbered content + file metadata) into a compact preview."""
is_read_file = (
isinstance(value, dict)
and isinstance(value.get("content"), str)
and all(k in value for k in ("total_lines", "file_size", "is_binary", "is_image"))
and not value.get("error")
)
if not is_read_file:
return value
normalized: dict[str, Any] = {}
if isinstance(args, dict):
if tool_name == "read_file" and isinstance(args, dict):
if isinstance(args.get("path"), str) and args["path"]:
normalized["path"] = args["path"]
normalized.update({key: args[key] for key in ("offset", "limit") if isinstance(args.get(key), int)})
lines = _parse_read_file_lines(value.get("content", ""))
content = value.get("content", "")
matches = [_READ_FILE_LINE_RE.match(raw) for raw in content.splitlines()] if isinstance(content, str) and content else []
lines = [{"line": int(m.group(1)), "text": m.group(2)} for m in matches] if matches and all(matches) else []
if lines:
normalized["returned_lines"] = {"start": lines[0]["line"], "end": lines[-1]["line"], "count": len(lines)}
head, tail = _READ_FILE_HEAD_LINES, _READ_FILE_TAIL_LINES
@@ -359,19 +346,6 @@ def _normalize_read_file_payload(value: dict[str, Any], *, args: Any = None) ->
return normalized
def _normalize_payload(value: Any, *, tool_name: str = "", args: Any = None) -> Any:
"""Collapse a read_file result (line-numbered content + file metadata) into a compact preview."""
is_read_file = (
isinstance(value, dict)
and isinstance(value.get("content"), str)
and all(k in value for k in ("total_lines", "file_size", "is_binary", "is_image"))
and not value.get("error")
)
if is_read_file:
return _normalize_read_file_payload(value, args=args if tool_name == "read_file" else None)
return value
def _safe_value(value: Any, *, max_chars: Optional[int] = None, depth: int = 0,
parse_json_strings: bool = False) -> Any:
max_chars = max_chars if max_chars is not None else int(_env("HERMES_LANGFUSE_MAX_CHARS", "12000") or "12000")
@@ -397,13 +371,6 @@ def _safe_value(value: Any, *, max_chars: Optional[int] = None, depth: int = 0,
return _truncate_text(repr(value), max_chars)
def _extract_last_user_message(messages: Any) -> Any:
if not isinstance(messages, list):
return None
last = next((m for m in reversed(messages) if isinstance(m, dict) and m.get("role") == "user"), None)
return None if last is None else {"role": "user", "content": _capture_content(last.get("content"))}
def _coerce_request_messages(*, request_messages: Any = None, messages: Any = None,
conversation_history: Any = None, user_message: Any = None) -> list[dict[str, Any]]:
for candidate in (request_messages, messages, conversation_history):
@@ -417,14 +384,10 @@ def _serialize_system_prompt(system_prompt: Any) -> Optional[dict[str, Any]]:
if isinstance(system_prompt, str):
text = system_prompt.strip()
elif isinstance(system_prompt, list):
parts: list[str] = []
for block in system_prompt:
if isinstance(block, dict):
# Anthropic: {"type": "text", "text": ...}; Bedrock Converse: {"text": ...}.
block = block.get("text", "") if block.get("type") in ("text", None) and "text" in block else None
if isinstance(block, str) and block:
parts.append(block)
text = "\n\n".join(parts)
# Anthropic: {"type": "text", "text": ...}; Bedrock Converse: {"text": ...}; or bare strings.
blocks = ((b.get("text", "") if b.get("type") in ("text", None) and "text" in b else None)
if isinstance(b, dict) else b for b in system_prompt)
text = "\n\n".join(b for b in blocks if isinstance(b, str) and b)
else:
return None
return {"role": "system", "content": _capture_content(text)} if text else None
@@ -432,52 +395,34 @@ def _serialize_system_prompt(system_prompt: Any) -> Optional[dict[str, Any]]:
def _messages_for_langfuse_input(*, request_messages: Any = None, messages: Any = None,
conversation_history: Any = None, user_message: Any = None,
system_prompt: Any = None, pre_coerced: Any = None) -> list[dict[str, Any]]:
"""Generation input, prepending ``system_prompt`` when the provider split it out
of messages. ``pre_coerced`` skips a second ``_coerce_request_messages``."""
raw = pre_coerced if pre_coerced is not None else _coerce_request_messages(
request_messages=request_messages, messages=messages,
conversation_history=conversation_history, user_message=user_message,
)
system_prompt: Any = None) -> list[dict[str, Any]]:
"""Generation input, prepending ``system_prompt`` when the provider split it out of messages."""
raw = _coerce_request_messages(request_messages=request_messages, messages=messages,
conversation_history=conversation_history, user_message=user_message)
system_msg = None if raw and raw[0].get("role") == "system" else _serialize_system_prompt(system_prompt)
serialized = _serialize_messages(raw)
return serialized if system_msg is None else [system_msg, *serialized]
def _serialize_message(message: dict[str, Any]) -> dict[str, Any]:
role, is_tool = message.get("role"), message.get("role") == "tool"
return {
"role": role, "content": _capture_content(message.get("content"), parse_json_strings=is_tool),
**({"tool_call_id": message["tool_call_id"]} if is_tool and message.get("tool_call_id") else {}),
**({"name": _safe_value(message["name"])} if is_tool and message.get("name") else {}),
**({"tool_calls": _capture_content(message["tool_calls"], parse_json_strings=True)} if message.get("tool_calls") else {}),
}
def _serialize_messages(messages: Any) -> list[dict[str, Any]]:
if not isinstance(messages, list):
return []
serialized = []
for message in messages[-12:]:
if not isinstance(message, dict):
continue
role = message.get("role")
item = {"role": role, "content": _capture_content(message.get("content"), parse_json_strings=(role == "tool"))}
if role == "tool":
if message.get("tool_call_id"):
item["tool_call_id"] = message["tool_call_id"]
if message.get("name"):
item["name"] = _safe_value(message["name"])
if message.get("tool_calls"):
item["tool_calls"] = _capture_content(message.get("tool_calls"), parse_json_strings=True)
serialized.append(item)
return serialized
return [_serialize_message(m) for m in messages[-12:] if isinstance(m, dict)] if isinstance(messages, list) else []
def _serialize_tool_calls(tool_calls: Any) -> list[dict[str, Any]]:
serialized = []
for tool_call in tool_calls or ():
fn = getattr(tool_call, "function", None)
name = getattr(fn, "name", None)
safe_arguments = _capture_content(getattr(fn, "arguments", None))
serialized.append({
"id": getattr(tool_call, "id", None),
"type": getattr(tool_call, "type", None) or "function",
"name": name,
"arguments": safe_arguments,
"function": {"name": name, "arguments": safe_arguments},
})
return serialized
def _serialize_tool_call(tool_call: Any) -> dict[str, Any]:
fn = getattr(tool_call, "function", None)
name, safe_arguments = getattr(fn, "name", None), _capture_content(getattr(fn, "arguments", None))
return {"id": getattr(tool_call, "id", None), "type": getattr(tool_call, "type", None) or "function",
"name": name, "arguments": safe_arguments, "function": {"name": name, "arguments": safe_arguments}}
def _serialize_assistant_message(message: Any) -> dict[str, Any]:
@@ -486,7 +431,7 @@ def _serialize_assistant_message(message: Any) -> dict[str, Any]:
return {
"content": _capture_content(getattr(message, "content", None)),
"reasoning": None if reasoning is None else _capture_content(reasoning),
"tool_calls": _serialize_tool_calls(getattr(message, "tool_calls", None)),
"tool_calls": [_serialize_tool_call(tc) for tc in getattr(message, "tool_calls", None) or ()],
}
@@ -540,32 +485,25 @@ def _canonical_usage_and_cost(canonical: Any, *, provider: str, model: str,
return usage_details, cost_details
def _usage_and_cost(response: Any, *, provider: str, api_mode: str, model: str, base_url: str) -> tuple[dict[str, int], dict[str, float]]:
def _usage_and_cost(response: Any, *, provider: str, model: str, base_url: str, api_mode: str = "",
usage: Optional[dict] = None) -> tuple[dict[str, int], dict[str, float]]:
"""Langfuse usage/cost maps from ``response.usage`` (post_llm_call) or, when ``usage``
is given (post_api_request), from that pre-built CanonicalUsage summary dict."""
raw_usage = getattr(response, "usage", None)
if not raw_usage:
if usage is None and not raw_usage:
return {}, {}
try:
from agent.usage_pricing import normalize_usage
from agent.usage_pricing import CanonicalUsage, normalize_usage
canonical = normalize_usage(raw_usage, provider=provider, api_mode=api_mode)
return _canonical_usage_and_cost(canonical, provider=provider, model=model, base_url=base_url)
except Exception as exc: # pragma: no cover - fail-open
_debug(f"usage normalization failed: {exc}")
return {}, {}
def _summary_usage_and_cost(usage: dict, *, provider: str, model: str, base_url: str) -> tuple[dict[str, int], dict[str, float]]:
"""post_api_request path: usage arrives as a pre-built CanonicalUsage summary dict."""
try:
from agent.usage_pricing import CanonicalUsage
canonical = CanonicalUsage(
canonical = normalize_usage(raw_usage, provider=provider, api_mode=api_mode) if usage is None else CanonicalUsage(
output_tokens=usage.get("output_tokens", 0) or usage.get("completion_tokens", 0),
request_count=usage.get("request_count", 1),
**{attr: usage.get(attr, 0) for attr in ("input_tokens", "cache_read_tokens", "cache_write_tokens", "reasoning_tokens")},
)
return _canonical_usage_and_cost(canonical, provider=provider, model=model, base_url=base_url)
except Exception:
except Exception as exc: # pragma: no cover - fail-open
if usage is None:
_debug(f"usage normalization failed: {exc}")
return {}, {}
@@ -573,7 +511,9 @@ def _start_root_trace(task_key: str, *, task_id: str, session_id: str, platform:
api_mode: str, messages: Any, client: Langfuse,
turn_id: str = "", api_request_id: str = "") -> TraceState:
trace_id = client.create_trace_id(seed=f"{session_id or 'sessionless'}::{task_id or task_key}")
trace_input = _extract_last_user_message(messages)
last_user = next((m for m in reversed(messages) if isinstance(m, dict) and m.get("role") == "user"), None) \
if isinstance(messages, list) else None
trace_input = None if last_user is None else {"role": "user", "content": _capture_content(last_user.get("content"))}
metadata = {
"source": "hermes", "task_id": task_id, "turn_id": turn_id, "api_request_id": api_request_id,
"platform": platform, "provider": provider, "model": model, "api_mode": api_mode,
@@ -598,19 +538,16 @@ def _start_root_trace(task_key: str, *, task_id: str, session_id: str, platform:
if root_ctx is None:
root_ctx, root_span = open_root()
# SDK v3 uses update_trace(); failures must never block the turn.
try:
with _failsafe("update_trace(input)"): # SDK v3 uses update_trace()
root_span.update_trace(input=trace_input)
except Exception as exc:
_debug(f"update_trace(input) failed: {exc}")
_debug(f"started trace {trace_id} for {task_key}")
return TraceState(trace_id=trace_id, root_ctx=root_ctx, root_span=root_span)
def _start_child_observation(state: TraceState, *, client: Langfuse, name: str, as_type: str,
input_value: Any, metadata: Optional[dict] = None,
model: Optional[str] = None, model_parameters: Optional[dict] = None) -> Any:
def _start_child_observation(state: TraceState, *, name: str, as_type: str, input_value: Any,
metadata: Optional[dict] = None, model: Optional[str] = None,
model_parameters: Optional[dict] = None) -> Any:
return state.root_span.start_observation(name=name, as_type=as_type, input=input_value, metadata=metadata or {},
model=model, model_parameters=model_parameters)
@@ -619,15 +556,13 @@ def _end_observation(observation: Any, *, output: Any = None, metadata: Optional
usage_details: Optional[dict] = None, cost_details: Optional[dict] = None) -> None:
if observation is None:
return
try:
update_kwargs: Dict[str, Any] = {} if output is None else {"output": output}
update_kwargs.update({k: v for k, v in (("metadata", metadata), ("usage_details", usage_details),
("cost_details", cost_details)) if v})
with _failsafe("end observation"):
update_kwargs = {**({} if output is None else {"output": output}),
**{k: v for k, v in (("metadata", metadata), ("usage_details", usage_details),
("cost_details", cost_details)) if v}}
if update_kwargs:
observation.update(**update_kwargs)
observation.end()
except Exception as exc: # pragma: no cover - fail-open
_debug(f"end observation failed: {exc}")
def _end_children(state: TraceState, *, include_subagents: bool = False) -> None:
@@ -637,45 +572,15 @@ def _end_children(state: TraceState, *, include_subagents: bool = False) -> None
_end_observation(observation)
def _exit_root_ctx(state: TraceState) -> None:
# Unwind the root context manager now, while opentelemetry.trace.Span is
# still a real type; GC-driven close at interpreter teardown raises
# TypeError inside use_span's isinstance check.
if state.root_ctx is not None:
try:
state.root_ctx.__exit__(None, None, None)
except Exception: # pragma: no cover - fail-open
pass
def _end_root(state: TraceState, label: str) -> None:
"""End the root span then unwind its context; never raises."""
try:
with _failsafe(label):
state.root_span.end()
_exit_root_ctx(state)
except Exception as exc: # pragma: no cover - fail-open
_debug(f"{label} failed: {exc}")
def _merge_trace_output(output: Any, state: TraceState) -> Any:
if not state.turn_tool_calls:
return output
merged = dict(output) if isinstance(output, dict) else {"content": output}
merged["tool_calls"] = list(state.turn_tool_calls)
return merged
def _evict_stale_locked() -> None:
"""Evict least-recently-updated state down to ``_MAX_TRACE_STATE - 1``. Caller
holds ``_STATE_LOCK`` and inserts exactly one entry afterwards; evicted roots
are ended so they don't dangle on the Langfuse side."""
over = len(_TRACE_STATE) - (_MAX_TRACE_STATE - 1)
if over <= 0:
return
stale = sorted(_TRACE_STATE.items(), key=lambda kv: kv[1].last_updated_at)[:over]
for key, state in stale:
_TRACE_STATE.pop(key, None)
_end_root(state, "evict stale trace")
# Unwind the root context manager now, while opentelemetry.trace.Span is
# still a real type; GC-driven close at interpreter teardown raises
# TypeError inside use_span's isinstance check.
if state.root_ctx is not None:
state.root_ctx.__exit__(None, None, None)
def _finalize_all_traces() -> None:
@@ -687,11 +592,8 @@ def _finalize_all_traces() -> None:
states = list(_TRACE_STATE.items())
_TRACE_STATE.clear()
for key, state in states:
try:
with _failsafe(f"atexit finalize for {key}"): # _end_root never raises
_end_children(state, include_subagents=True)
except Exception as exc: # pragma: no cover - fail-open
_debug(f"atexit finalize failed for {key}: {exc}")
else:
_end_root(state, f"atexit finalize for {key}")
if states:
_flush(_get_langfuse())
@@ -699,41 +601,34 @@ def _finalize_all_traces() -> None:
def _flush(client: Any) -> None:
if client is not None:
try:
with contextlib.suppress(Exception):
client.flush()
except Exception:
pass
def _finish_trace(task_key: str, *, output: Any = None) -> None:
client = _get_langfuse()
if client is None:
return
with _STATE_LOCK:
state = _TRACE_STATE.pop(task_key, None)
state = _TRACE_STATE.pop(task_key, None) if client is not None else None
if state is None:
return
try:
_end_children(state)
final_output = _merge_trace_output(output, state)
final_output = output
if state.turn_tool_calls:
final_output = dict(output) if isinstance(output, dict) else {"content": output}
final_output["tool_calls"] = list(state.turn_tool_calls)
if final_output is not None:
# update_trace sets TRACE-level I/O (SDK v3); root I/O via update().
# Neither may prevent end(), else children export without a root.
for method, label in (("update_trace", "update_trace(output)"), ("update", "root update(output)")):
try:
with _failsafe(label):
getattr(state.root_span, method)(output=final_output)
except Exception as exc:
_debug(f"{label} failed: {exc}")
_end_root(state, "root end()")
except Exception as exc: # pragma: no cover - fail-open
_debug(f"finish trace failed: {exc}")
# Last-chance end so an unexpected error still exports the root.
try:
with contextlib.suppress(Exception): # last-chance end so the root still exports
state.root_span.end()
except Exception:
pass
finally:
_flush(client)
@@ -750,9 +645,8 @@ def _client_and_key(task_id: str, session_id: str, turn_id: str, api_request_id:
return client, _trace_key(task_id, session_id, turn_id=turn_id, api_request_id=api_request_id)
def _add_duration(metadata: Dict[str, Any], api_duration: Any) -> None:
if api_duration and api_duration > 0:
metadata["api_duration_s"] = round(api_duration, 3)
def _duration_meta(api_duration: Any) -> Dict[str, Any]:
return {"api_duration_s": round(api_duration, 3)} if api_duration and api_duration > 0 else {}
def _pop_generation(task_key: str, api_call_count: Any) -> tuple[Optional[TraceState], Any]:
@@ -763,11 +657,16 @@ def _pop_generation(task_key: str, api_call_count: Any) -> tuple[Optional[TraceS
def _get_or_start_state_locked(task_key: str, **root_kwargs: Any) -> TraceState:
"""Caller must hold ``_STATE_LOCK``. Starts a root trace if the key is new."""
"""Caller must hold ``_STATE_LOCK``. Starts a root trace if the key is new, first
evicting least-recently-updated state down to ``_MAX_TRACE_STATE - 1`` (evicted
roots are ended so they don't dangle on the Langfuse side)."""
state = _TRACE_STATE.get(task_key)
if state is None:
state = _start_root_trace(task_key, **root_kwargs)
_evict_stale_locked()
over = len(_TRACE_STATE) - (_MAX_TRACE_STATE - 1)
for key, stale in sorted(_TRACE_STATE.items(), key=lambda kv: kv[1].last_updated_at)[:max(over, 0)]:
_TRACE_STATE.pop(key, None)
_end_root(stale, "evict stale trace")
_TRACE_STATE[task_key] = state
state.last_updated_at = time.time()
return state
@@ -784,11 +683,8 @@ def on_pre_llm_call(*, task_id: str = "", session_id: str = "", platform: str =
if client is None:
return
with _STATE_LOCK:
_get_or_start_state_locked(
task_key, task_id=task_id, session_id=session_id, platform=platform, provider=provider,
model=model, api_mode=api_mode, messages=messages, client=client,
turn_id=turn_id, api_request_id=api_request_id,
)
_get_or_start_state_locked(task_key, task_id=task_id, session_id=session_id, platform=platform, provider=provider, model=model,
api_mode=api_mode, messages=messages, client=client, turn_id=turn_id, api_request_id=api_request_id)
def _emit_moa_reference_generations(state: TraceState, *, client: Langfuse, references: Any) -> None:
@@ -816,17 +712,13 @@ def _emit_moa_reference_generations(state: TraceState, *, client: Langfuse, refe
cost_details = {"total": float(cost_usd)} if isinstance(cost_usd, (int, float)) else {}
label = ref.get("label") or "advisor"
metadata = {"moa_role": "reference", "label": label}
metadata.update({k: ref[k] for k in ("provider", "cost_status", "cost_source", "temperature") if ref.get(k) is not None})
metadata = {"moa_role": "reference", "label": label,
**{k: ref[k] for k in ("provider", "cost_status", "cost_source", "temperature") if ref.get(k) is not None}}
observation = _start_child_observation(
state, client=client, name=f"MoA advisor: {label}", as_type="generation",
input_value=None, metadata=metadata, model=ref.get("model"),
)
_end_observation(
observation, output=_capture_content(ref.get("output")),
usage_details=usage_details, cost_details=cost_details, metadata=metadata,
)
observation = _start_child_observation(state, name=f"MoA advisor: {label}", as_type="generation", input_value=None,
metadata=metadata, model=ref.get("model"))
_end_observation(observation, output=_capture_content(ref.get("output")), usage_details=usage_details,
cost_details=cost_details, metadata=metadata)
def on_pre_llm_request(*, task_id: str = "", session_id: str = "", platform: str = "", model: str = "",
@@ -845,32 +737,27 @@ def on_pre_llm_request(*, task_id: str = "", session_id: str = "", platform: str
if isinstance(body_model, str) and body_model:
model = body_model
input_messages = _coerce_request_messages(
request_messages=request_messages, messages=messages,
conversation_history=conversation_history, user_message=user_message,
)
langfuse_input = _messages_for_langfuse_input(system_prompt=system_prompt, pre_coerced=input_messages)
input_messages = _coerce_request_messages(request_messages=request_messages, messages=messages,
conversation_history=conversation_history, user_message=user_message)
langfuse_input = _messages_for_langfuse_input(request_messages=input_messages, system_prompt=system_prompt)
has_system = bool(langfuse_input) and langfuse_input[0].get("role") == "system"
system_chars = len(str(langfuse_input[0].get("content") or "")) if has_system else 0
req_key = _request_key(api_call_count)
with _STATE_LOCK:
state = _get_or_start_state_locked(
task_key, task_id=task_id, session_id=session_id, platform=platform, provider=provider,
model=model, api_mode=api_mode, messages=input_messages, client=client,
turn_id=turn_id, api_request_id=api_request_id,
)
task_key, task_id=task_id, session_id=session_id, platform=platform, provider=provider, model=model,
api_mode=api_mode, messages=input_messages, client=client, turn_id=turn_id, api_request_id=api_request_id)
previous = state.generations.pop(req_key, None)
if previous is not None:
_end_observation(previous)
gen_metadata = {
"provider": provider, "platform": platform, "api_mode": api_mode, "base_url": base_url,
"message_count": message_count, "approx_input_tokens": approx_input_tokens,
**({"system_prompt_chars": system_chars} if system_chars else {}),
}
if system_chars:
gen_metadata["system_prompt_chars"] = system_chars
state.generations[req_key] = _start_child_observation(
state, client=client, name=f"LLM call {api_call_count}", as_type="generation",
state, name=f"LLM call {api_call_count}", as_type="generation",
input_value=langfuse_input, metadata=gen_metadata, model=model,
model_parameters={"api_mode": api_mode, "provider": provider},
)
@@ -904,11 +791,8 @@ def on_post_llm_call(*, task_id: str = "", session_id: str = "", provider: str =
elif assistant_response is not None:
output = {"content": _capture_content(assistant_response), "reasoning": None, "tool_calls": []}
else:
output = {
"content": f"[{assistant_content_chars} chars]" if assistant_content_chars else None,
"reasoning": None,
"tool_calls": [{"id": f"tc_{i}"} for i in range(assistant_tool_call_count or 0)],
}
output = {"content": f"[{assistant_content_chars} chars]" if assistant_content_chars else None, "reasoning": None,
"tool_calls": [{"id": f"tc_{i}"} for i in range(assistant_tool_call_count or 0)]}
if output.get("tool_calls"):
state.turn_tool_calls.extend(output["tool_calls"])
@@ -916,22 +800,15 @@ def on_post_llm_call(*, task_id: str = "", session_id: str = "", provider: str =
# post_api_request's ``response`` is a sanitized dict with no ``.usage``;
# gate on the attribute so the usage-dict fallback is actually reached.
if getattr(response, "usage", None) is not None:
usage_details, cost_details = _usage_and_cost(
response, provider=provider, api_mode=api_mode, model=model, base_url=base_url,
)
usage_details, cost_details = _usage_and_cost(response, provider=provider, api_mode=api_mode, model=model, base_url=base_url)
elif isinstance(usage, dict) and usage:
usage_details, cost_details = _summary_usage_and_cost(
usage, provider=provider, model=model, base_url=base_url,
)
usage_details, cost_details = _usage_and_cost(None, provider=provider, model=model, base_url=base_url, usage=usage)
else:
usage_details, cost_details = {}, {}
gen_metadata: Dict[str, Any] = {"tool_call_count": len(output.get("tool_calls", [])) or assistant_tool_call_count}
_add_duration(gen_metadata, api_duration)
if finish_reason:
gen_metadata["finish_reason"] = finish_reason
_end_observation(generation, output=output, usage_details=usage_details,
cost_details=cost_details, metadata=gen_metadata)
gen_metadata = {"tool_call_count": len(output.get("tool_calls", [])) or assistant_tool_call_count,
**_duration_meta(api_duration), **({"finish_reason": finish_reason} if finish_reason else {})}
_end_observation(generation, output=output, usage_details=usage_details, cost_details=cost_details, metadata=gen_metadata)
has_tools = bool(getattr(assistant_message, "tool_calls", None)) if assistant_message else assistant_tool_call_count > 0
if not has_tools and output.get("content"):
@@ -948,11 +825,8 @@ def on_pre_tool_call(*, tool_name: str = "", args: Any = None, task_id: str = ""
state = _TRACE_STATE.get(task_key)
if state is None:
return
observation = _start_child_observation(
state, client=client, name=f"Tool: {tool_name}", as_type="tool",
input_value=_capture_content(args),
metadata={"tool_name": tool_name, "tool_call_id": tool_call_id},
)
observation = _start_child_observation(state, name=f"Tool: {tool_name}", as_type="tool", input_value=_capture_content(args),
metadata={"tool_name": tool_name, "tool_call_id": tool_call_id})
if tool_call_id:
state.tools[tool_call_id] = observation
else:
@@ -976,7 +850,7 @@ def on_post_tool_call(*, tool_name: str = "", args: Any = None, result: Any = No
if observation is None:
return
safe_result_value = _capture_tool_result(result, tool_name=tool_name, args=args)
safe_result_value = _capture_content(result, tool_result_of=(tool_name, args))
# Backfill so the generation's tool_call record carries the result alongside arguments.
if tool_call_id:
@@ -984,15 +858,12 @@ def on_post_tool_call(*, tool_name: str = "", args: Any = None, result: Any = No
state = _TRACE_STATE.get(task_key)
calls = state.turn_tool_calls if state is not None else []
tool_call = next((tc for tc in reversed(calls) if tc.get("id") == tool_call_id), None)
if tool_call is not None:
tool_call["output"] = safe_result_value
if isinstance(tool_call.get("function"), dict):
tool_call["function"]["output"] = safe_result_value
for target in (tool_call, tool_call.get("function")) if tool_call is not None else ():
if isinstance(target, dict):
target["output"] = safe_result_value
_end_observation(
observation, output=safe_result_value,
metadata={"tool_name": tool_name, "args": _capture_content(args, parse_json_strings=True)},
)
_end_observation(observation, output=safe_result_value,
metadata={"tool_name": tool_name, "args": _capture_content(args, parse_json_strings=True)})
def on_api_request_error(*, task_id: str = "", session_id: str = "", api_call_count: int = 0,
@@ -1013,18 +884,16 @@ def on_api_request_error(*, task_id: str = "", session_id: str = "", api_call_co
error_type, error_message = str(error.get("type") or ""), str(error.get("message") or "")
# Error messages can embed request fragments (URLs w/ keys, prompt echoes) — capture-pipeline them.
error_metadata: Dict[str, Any] = {"error": True, "error_type": error_type, "error_message": _capture_content(error_message)}
error_metadata.update({k: v for k, v in (("status_code", status_code), ("retry_count", retry_count),
("max_retries", max_retries), ("retryable", retryable)) if v is not None})
if reason:
error_metadata["reason"] = str(reason)
_add_duration(error_metadata, api_duration)
error_metadata: Dict[str, Any] = {
"error": True, "error_type": error_type, "error_message": _capture_content(error_message),
**{k: v for k, v in (("status_code", status_code), ("retry_count", retry_count), ("max_retries", max_retries),
("retryable", retryable), ("reason", str(reason) if reason else None)) if v is not None},
**_duration_meta(api_duration),
}
if generation is not None:
try:
with _failsafe("error-level update"):
generation.update(level="ERROR", status_message=(error_type or "api_request_error")[:200])
except Exception as exc: # pragma: no cover - fail-open
_debug(f"error-level update failed: {exc}")
_end_observation(generation, metadata=error_metadata)
# A retryable failure is followed by another pre_api_request on the same
@@ -1046,12 +915,9 @@ def on_session_finalize(*, session_id: str = "", reason: str = "", **_: Any) ->
# This session's traces (all, when no session_id). Keys carry the session as
# "session:<id>" or "task:<id>" (gateway: task_id == session_id) or bare legacy id.
fragments = (f"session:{session_id}", f"task:{session_id}")
with _STATE_LOCK:
if session_id:
fragments = (f"session:{session_id}", f"task:{session_id}")
keys = [k for k in _TRACE_STATE if k == session_id or any(f in k for f in fragments)]
else:
keys = list(_TRACE_STATE)
keys = [k for k in _TRACE_STATE if not session_id or k == session_id or any(f in k for f in fragments)]
for key in keys:
_finish_trace(key)
_flush(client)
@@ -1060,10 +926,8 @@ def on_session_finalize(*, session_id: str = "", reason: str = "", **_: Any) ->
# cached client must keep exporting). Doing it while modules are intact keeps
# the SDK's atexit handler off torn-down opentelemetry globals (TypeError on quit).
if reason == "shutdown" and callable(getattr(client, "shutdown", None)):
try:
with _failsafe("langfuse shutdown"):
client.shutdown()
except Exception as exc: # pragma: no cover - fail-open
_debug(f"langfuse shutdown failed: {exc}")
def on_subagent_start(*, parent_turn_id: str = "", parent_subagent_id: Any = None,
@@ -1077,13 +941,11 @@ def on_subagent_start(*, parent_turn_id: str = "", parent_subagent_id: Any = Non
state = _state_for_turn(parent_turn_id)
if state is None:
return
metadata = {"child_session_id": child_session_id, "child_subagent_id": child_subagent_id, "child_role": child_role}
if parent_subagent_id:
metadata["parent_subagent_id"] = parent_subagent_id
metadata = {"child_session_id": child_session_id, "child_subagent_id": child_subagent_id, "child_role": child_role,
**({"parent_subagent_id": parent_subagent_id} if parent_subagent_id else {})}
state.subagents[str(child_session_id)] = _start_child_observation(
state, client=client, name=f"Subagent: {child_role or 'delegate'}", as_type="span",
input_value=_capture_content(child_goal), metadata=metadata,
)
state, name=f"Subagent: {child_role or 'delegate'}", as_type="span",
input_value=_capture_content(child_goal), metadata=metadata)
def on_subagent_stop(*, parent_turn_id: str = "", child_session_id: Any = None, child_role: str = "",
@@ -1100,11 +962,9 @@ def on_subagent_stop(*, parent_turn_id: str = "", child_session_id: Any = None,
if observation is None:
return
metadata: Dict[str, Any] = {"child_role": child_role}
metadata.update({k: v for k, v in (("status", child_status), ("duration_ms", duration_ms)) if v})
if isinstance(tool_call_history, list):
metadata["tool_call_count"] = len(tool_call_history)
metadata["tool_calls"] = _capture_content(tool_call_history)
metadata = {"child_role": child_role, **{k: v for k, v in (("status", child_status), ("duration_ms", duration_ms)) if v},
**({"tool_call_count": len(tool_call_history), "tool_calls": _capture_content(tool_call_history)}
if isinstance(tool_call_history, list) else {})}
_end_observation(observation, output=_capture_content(child_summary), metadata=metadata)