scugnizz-llama-training / agent_loop.py
Reverendo's picture
Upload agent_loop.py with huggingface_hub
102ba2a verified
Raw History Blame Contribute Delete
3.87 kB
#!/usr/bin/env python3
"""Hermes tool loop: generate → parse tool_call → execute → tool_response → repeat."""
from __future__ import annotations
import json
from typing import Any, Callable
from hermes_agent_eval import build_prompt, hermes_tool_content, tool_response
from web_search_tool import (
parse_tool_calls,
tool_names_from_schema,
truncate_at_first_tool_call,
unknown_tool_result,
)
def _call_key(name: str, arguments: dict) -> str:
try:
return f"{name}:{json.dumps(arguments or {}, sort_keys=True)}"
except TypeError:
return f"{name}:{arguments!r}"
def _tool_ok(result: Any) -> bool:
if not isinstance(result, dict):
return True
if result.get("error") and result.get("success") is False:
return False
if result.get("error") and not (
result.get("results")
or ((result.get("data") or {}).get("web"))
):
return False
if "results" in result:
return True # empty results[] is valid Hermes content (refuse-after-empty)
if isinstance((result.get("data") or {}).get("web"), list):
return True # hermes-agent web_search envelope
if "hits" in result:
return True
if "result" in result:
return True
if result.get("ok") or result.get("content") is not None:
return True
return True
def run_agent_loop(
generate_fn: Callable[[str], str],
tok,
tools: list,
turns: list[dict],
execute_fn: Callable[[str, dict], Any],
*,
max_rounds: int = 4,
user_question: str = "",
) -> dict:
"""Mutates `turns` in place. Returns {final, turns, trace, tool_rounds}.
Thin loop only — no answer nudges / forced prose. Train the model for that.
Exact same successful call twice stops the loop (Space cost guard only).
"""
allowed = tool_names_from_schema(tools)
trace: list[dict] = []
tool_round = 0
final = ""
last_ok_key: str | None = None
while True:
prompt = build_prompt(tok, tools, turns)
output = truncate_at_first_tool_call(generate_fn(prompt))
final = output
trace.append({"role": "assistant", "content": output})
calls = parse_tool_calls(output)
if not calls:
return {
"final": output,
"turns": turns,
"trace": trace,
"tool_rounds": tool_round,
}
if tool_round >= max_rounds:
return {
"final": output,
"turns": turns,
"trace": trace,
"tool_rounds": tool_round,
"stopped": "max_tool_rounds",
}
call = calls[0]
name = call["name"]
arguments = call.get("arguments") or {}
key = _call_key(name, arguments)
# Exact repeat of a successful call → stop (prevents infinite billable loops).
if last_ok_key is not None and key == last_ok_key:
print(f"call-loop blocked: exact repeat {name}", flush=True)
return {
"final": output,
"turns": turns,
"trace": trace,
"tool_rounds": tool_round,
"stopped": "call_loop",
}
if name not in allowed:
result = hermes_tool_content(name, unknown_tool_result(name, allowed))
else:
result = execute_fn(name, arguments)
turns.append({"from": "gpt", "value": output})
turns.append(
{
"from": "tool",
"value": tool_response(f"{name}:{tool_round}", name, result),
}
)
trace.append(
{"role": "tool", "name": name, "arguments": arguments, "content": result}
)
if _tool_ok(result):
last_ok_key = key
tool_round += 1