Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions agent_core/loop_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,18 @@ class LoopConfig:
system_addendum_per_call: str | Callable[[], str] = ""
system_addendum_min_turn: int = 0
system_addendum_per_call_role: str = "system"
# Keyword-only to preserve every existing positional LoopConfig argument.
# When unset, retain the legacy stream_llm_tokens/watchdog selection.
stream_transport: bool | None = field(default=None, kw_only=True)
# How strictly native tool-call arguments are checked before execution.
# "structural" (default): decodable JSON object with every top-level
# required property; blank arguments count as {} and property types are
# left to the tool. "strict": full JSON Schema. "off": no checks beyond
# the legacy empty-required-arguments retry. Text-mode calls are never
# blocked.
tool_argument_validation: Literal["structural", "strict", "off"] = field(
default="structural", kw_only=True,
)


@dataclass
Expand Down
131 changes: 90 additions & 41 deletions agent_core/runtime/loop/_call.py

Large diffs are not rendered by default.

69 changes: 68 additions & 1 deletion agent_core/runtime/loop/agent_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,10 @@
MultiFormatToolCallParser,
ToolCallParser,
)
from agent_core.runtime.loop.tool_call_validation import (
IDENTITY_REASONS,
native_tool_call_issues,
)
from agent_core.runtime.loop.tool_exec import (
DefaultToolResultPostProcessor,
ToolExecutionHooks,
Expand Down Expand Up @@ -454,7 +458,9 @@ async def _run_loop_inner(
for observer in obs
)
protocol = getattr(profile, "protocol", "chat_completions")
if cfg.stream_llm_tokens is not None:
if cfg.stream_transport is not None:
stream_llm_tokens = cfg.stream_transport
elif cfg.stream_llm_tokens is not None:
# The host stated it; the transport may be the reason (see LoopConfig).
stream_llm_tokens = cfg.stream_llm_tokens
else:
Expand Down Expand Up @@ -639,6 +645,7 @@ async def _run_loop_inner(
result_max_chars=tool_result_cap,
profile=profile,
max_images_in_history=policy.max_images_in_history,
native_tool_calls=getattr(response, "tool_calls", None),
)
total_tool_calls += tool_calls_executed
if stop_reason:
Expand Down Expand Up @@ -894,6 +901,7 @@ async def _on_delta(
response = await call_llm(
llm_for_turn, messages_for_call, cfg.llm_timeout, cfg.max_llm_retries, turn,
on_delta=_on_delta if stream_llm_tokens else None,
stream=stream_llm_tokens,
retry_wait_fixed=cfg.retry_wait_fixed,
runaway_state=metadata.setdefault(RUNAWAY_STATE_KEY, {}),
first_chunk_s=cfg.first_chunk_timeout,
Expand All @@ -908,6 +916,7 @@ async def _on_delta(
wall_deadline_remaining=runtime_hooks.wall_deadline_remaining,
chain_fallback_active=runtime_hooks.chain_fallback_active,
empty_completion_max_retries=cfg.empty_completion_max_retries,
tool_argument_validation=getattr(cfg, "tool_argument_validation", "structural"),
)
except LLMCallExhausted as exhausted:
if exhausted.reason == "empty_completion":
Expand Down Expand Up @@ -1433,6 +1442,37 @@ def _effective_tool_result_cap(
return cap if cap and cap > 0 else None


def _invalid_native_calls_by_id(
native_tool_calls: list[dict] | None, tool_map: dict[str, ToolLike], mode: str,
) -> dict[str, dict[str, Any]]:
"""Argument diagnostics for native calls, keyed by tool-call id."""
if not native_tool_calls or mode == "off":
return {}
schemas: dict[str, dict[str, Any] | None] = {}

def schema_for(name: str) -> dict[str, Any] | None:
if name not in schemas:
schemas[name] = None
to_schema = getattr(tool_map.get(name), "to_openai_schema", None)
if callable(to_schema):
try:
schema = to_schema()
except Exception:
logger.debug("to_openai_schema failed for tool %s", name, exc_info=True)
schema = None
function = schema.get("function") if isinstance(schema, dict) else None
parameters = function.get("parameters") if isinstance(function, dict) else None
if isinstance(parameters, dict):
schemas[name] = parameters
return schemas[name]

return {
str(issue["id"]): issue
for issue in native_tool_call_issues(native_tool_calls, schema_for, mode)
if issue["id"] and issue["reason"] not in IDENTITY_REASONS
}


async def _execute_tool_calls(
cfg: LoopConfig, obs: list, tool_map: dict[str, ToolLike], messages: list[Message], metadata: dict[str, Any],
turn: int, total_tool_calls: int, ctx: TurnContext, parsed_calls: list[dict],
Expand All @@ -1448,9 +1488,15 @@ async def _execute_tool_calls(
# -1 disables eviction. Not 0: a defaulted caller must not silently mean
# "throw every image away", which is what a 0 default would spell.
max_images_in_history: int = -1,
# The raw ``response.tool_calls``. Only these provider-native calls are
# checked before dispatch; text-mode calls keep their lenient handling.
native_tool_calls: list[dict] | None = None,
) -> tuple[str, int]:
executable: list[tuple[int, dict]] = []
synthetic: list[tuple[int, ToolResult]] = []
invalid_native = _invalid_native_calls_by_id(
native_tool_calls, tool_map, getattr(cfg, "tool_argument_validation", "structural"),
)
for idx, tc in enumerate(parsed_calls):
# Text-mode calls arrive without a provider id. Assign it here, once,
# off the ``parsed_calls`` index -- before the batch is split into
Expand All @@ -1465,6 +1511,27 @@ async def _execute_tool_calls(
metadata.update(tcv.metadata_updates)
if tcv.rewrite_args is not None:
tc = {**tc, "args": tcv.rewrite_args}
issue = None if tcv.rewrite_args is not None else invalid_native.get(str(tc.get("id") or ""))
invalid_reason = issue["reason"] if issue is not None else None
if invalid_reason:
recorded = metadata.setdefault("invalid_tool_calls", [])
if isinstance(recorded, list):
recorded.append({
"id": tc["id"], "name": tc.get("name"),
"raw_arguments": issue.get("raw_arguments") if issue else None,
"reason": invalid_reason,
})
invalid_args = tc.get("args")
synthetic.append((idx, ToolResult(
name=str(tc.get("name", "") or ""),
args=invalid_args if isinstance(invalid_args, dict) else {},
result=f"[invalid tool call] {invalid_reason}; re-issue the call with valid arguments.",
duration_ms=0,
tool_call_id=str(tc["id"]),
is_error=True,
error_kind="invalid_arguments",
)))
continue
if tcv.skip_with_result is not None:
raw_args = tc.get("args")
synthetic.append((idx, ToolResult(
Expand Down
237 changes: 237 additions & 0 deletions agent_core/runtime/loop/tool_call_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,237 @@
# pyright: reportUnknownVariableType=false, reportUnknownMemberType=false, reportUnknownArgumentType=false, reportMissingModuleSource=false
"""Validate assembled native tool calls before any tool side effect.

Three levels, selected by ``LoopConfig.tool_argument_validation``:

``"structural"`` (default)
Arguments must decode to a JSON object that carries every top-level
``required`` property. A blank argument string is ``{}`` -- several
providers (Anthropic streaming among them) encode a zero-argument call
that way. Property types are *not* checked, so tools that coerce
``"5"`` to ``5`` keep working exactly as before.
``"strict"``
The structural checks plus full JSON Schema validation.
``"off"``
No validation; only the legacy empty-required-arguments retry remains.

A schema that is itself invalid never fails a call: validation for that tool
is skipped with a warning, matching how providers accept loose schemas.
"""

from __future__ import annotations

import json
import logging
import uuid
from collections.abc import Callable, Mapping
from functools import lru_cache
from typing import Any, Literal

logger = logging.getLogger(__name__)

ToolArgumentValidation = Literal["structural", "strict", "off"]
TOOL_ARGUMENT_VALIDATION_MODES: frozenset[str] = frozenset({"structural", "strict", "off"})

# Reasons that describe the call's identity rather than its arguments. They are
# reported, but a model retry cannot be expected to fix them: missing ids are
# repaired locally and nameless slots are already answered as dropped calls.
IDENTITY_REASONS: frozenset[str] = frozenset({"missing tool name", "missing tool call id"})


def _normalize_mode(mode: str | None) -> str:
if mode in TOOL_ARGUMENT_VALIDATION_MODES:
return str(mode)
if mode is not None:
logger.warning("Unknown tool_argument_validation %r; using 'structural'", mode)
return "structural"


def bound_tool_parameters(llm: Any) -> dict[str, dict[str, Any]]:
"""Map each tool bound on ``llm`` (OpenAI schema shape) to its parameters."""
parameters_by_name: dict[str, dict[str, Any]] = {}
for tool in getattr(llm, "tools", None) or []:
if not isinstance(tool, dict):
continue
function = tool.get("function")
if not isinstance(function, dict):
continue
name = function.get("name")
parameters = function.get("parameters")
if isinstance(name, str) and isinstance(parameters, dict):
parameters_by_name[name] = parameters
return parameters_by_name


@lru_cache(maxsize=256)
def _cached_validator(schema_json: str) -> Any | None:
from jsonschema.exceptions import SchemaError
from jsonschema.validators import validator_for

schema = json.loads(schema_json)
try:
validator_cls = validator_for(schema)
validator_cls.check_schema(schema)
return validator_cls(schema)
except SchemaError as exc:
logger.warning(
"Tool schema is not valid JSON Schema; skipping strict argument "
"validation for it: %s", exc.message,
)
return None


def _validator_for(schema: dict[str, Any]) -> Any | None:
try:
schema_json = json.dumps(schema, sort_keys=True, default=str)
except (TypeError, ValueError):
return None
return _cached_validator(schema_json)


def validate_arguments(value: Any, schema: dict[str, Any], path: str = "arguments") -> str | None:
"""Validate against the tool's complete JSON Schema.

Returns a diagnostic string, or ``None`` when the value is valid *or* the
schema itself cannot be used for validation.
"""
validator = _validator_for(schema)
if validator is None:
return None
from jsonschema.exceptions import best_match

try:
error = best_match(validator.iter_errors(value))
except Exception as exc: # unresolvable $ref and similar schema faults
logger.warning("Strict tool argument validation skipped: %s", exc)
return None
if error is None:
return None
location = ".".join(str(part) for part in error.absolute_path)
return f"{path}{'.' + location if location else ''}: {error.message}"


def _missing_required(args: dict[str, Any], schema: dict[str, Any]) -> str | None:
required = schema.get("required")
if not isinstance(required, list):
return None
missing = [field for field in required if isinstance(field, str) and field not in args]
if not missing:
return None
if len(missing) == 1:
return f"arguments: '{missing[0]}' is a required property"
return f"arguments: required properties missing: {', '.join(repr(m) for m in missing)}"


def argument_issue(raw: Any, schema: dict[str, Any] | None, mode: str = "structural") -> str | None:
"""Return why ``raw`` (a native ``function.arguments`` value) is invalid."""
mode = _normalize_mode(mode)
if mode == "off":
return None
if raw is None or (isinstance(raw, str) and not raw.strip()):
args: Any = {}
elif isinstance(raw, str):
try:
args = json.loads(raw)
except ValueError:
return "arguments are invalid JSON"
elif isinstance(raw, dict):
args = raw
else:
return "arguments must be a JSON object"
if not isinstance(args, dict):
return "arguments must be a JSON object"
if schema is None:
return None
missing = _missing_required(args, schema)
if missing:
return missing
if mode == "strict":
return validate_arguments(args, schema)
return None


def native_tool_call_issues(
tool_calls: Any,
schema_for: Callable[[str], dict[str, Any] | None] | Mapping[str, dict[str, Any]],
mode: str = "structural",
) -> list[dict[str, Any]]:
"""Diagnose native OpenAI-shape tool calls, keeping the raw arguments."""
lookup: Callable[[str], dict[str, Any] | None] = (
schema_for.get if isinstance(schema_for, Mapping) else schema_for
)
issues: list[dict[str, Any]] = []
for index, call in enumerate(tool_calls or []):
if not isinstance(call, dict):
continue
function = call.get("function")
if not isinstance(function, dict):
continue
name = function.get("name")
raw = function.get("arguments")
call_id = call.get("id")
if not isinstance(name, str) or not name:
reason: str | None = "missing tool name"
elif not isinstance(call_id, str) or not call_id:
reason = "missing tool call id"
else:
reason = argument_issue(raw, lookup(name), mode)
if reason:
issues.append({
"index": index, "id": call_id, "name": name,
"raw_arguments": raw, "reason": reason,
})
return issues


def invalid_native_tool_calls(
response: Any, llm: Any, mode: str = "structural",
) -> list[dict[str, Any]]:
"""Return argument diagnostics for ``response`` against ``llm``'s tools.

Identity problems (missing name / id) are excluded: a retry is not the
remedy for them -- see :func:`ensure_tool_call_ids`.
"""
return [
issue for issue in native_tool_call_issues(
getattr(response, "tool_calls", None), bound_tool_parameters(llm), mode,
)
if issue["reason"] not in IDENTITY_REASONS
]


def ensure_tool_call_ids(response: Any) -> list[int]:
"""Give every named native call without an id a unique one, in place.

The id lands on ``response.tool_calls`` before the assistant turn is
written to history, so the tool reply and the replayed assistant message
agree. Returns the indexes that were filled.
"""
filled: list[int] = []
for index, call in enumerate(getattr(response, "tool_calls", None) or []):
if not isinstance(call, dict):
continue
function = call.get("function")
name = function.get("name") if isinstance(function, dict) else call.get("name")
if not name:
continue
call_id = call.get("id")
if isinstance(call_id, str) and call_id:
continue
call["id"] = f"call_{uuid.uuid4().hex[:24]}"
filled.append(index)
if filled:
logger.warning("Assigned ids to %d native tool call(s) the provider sent without one", len(filled))
return filled


__all__ = [
"IDENTITY_REASONS",
"TOOL_ARGUMENT_VALIDATION_MODES",
"ToolArgumentValidation",
"argument_issue",
"bound_tool_parameters",
"ensure_tool_call_ids",
"invalid_native_tool_calls",
"native_tool_call_issues",
"validate_arguments",
]
1 change: 1 addition & 0 deletions changes/73.feature.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Streaming can now be selected independently of token observers (`LoopConfig.stream_transport`, `call_llm(stream=...)`). Native tool calls are checked before execution according to the new `LoopConfig.tool_argument_validation` (`"structural"` by default, `"strict"` for full JSON Schema, `"off"`). Invalid streamed calls are retried once using streaming, and invalid native calls that remain return an explicit tool error instead of running. Blank arguments count as `{}`, text-mode calls and tool-side type coercion are unaffected, malformed tool schemas never fail a call, and native calls missing an id get a generated one instead of breaking history replay.
Loading
Loading