diff --git a/agentic_loop/controller.py b/agentic_loop/controller.py index 2f242cc..ca2a22e 100644 --- a/agentic_loop/controller.py +++ b/agentic_loop/controller.py @@ -78,6 +78,7 @@ def run( enabled_rule_keys: list[str] | tuple[str, ...] | set[str] | None = None, disabled_rule_keys: list[str] | tuple[str, ...] | set[str] | None = None, workflow_key: str | None = None, + cancel_event=None, ) -> AgentResult: run_id = run_id or uuid.uuid4().hex state = AgentState(goal=goal) @@ -134,6 +135,18 @@ def run( ) for index in range(1, run_max_steps + 1): + # Cooperative cancellation (e.g. voice barge-in): stop at the step + # boundary so no further model calls or tools fire on an abandoned turn. + if cancel_event is not None and cancel_event.is_set(): + state.final_answer = "(interrupted)" + state.add_step(AgentStep(index=index, action="cancelled")) + self.logger.log("cancelled", {"step_index": index}, + session_id=session_id, run_id=run_id) + return self._result( + state=state, final_answer="(interrupted)", run_id=run_id, + session_id=session_id, workflow_key=workflow.key if workflow else None, + active_rule_keys=active_rule_keys_for_learning, matched_skills=matched_skills, + ) active_rules = self.rule_resolver.active_rules( session_id=session_id, run_id=run_id, @@ -294,6 +307,18 @@ def run( matched_skills=matched_skills, ) + # Don't fire a (possibly side-effecting) tool if we were cancelled + # while the model was producing this call. + if cancel_event is not None and cancel_event.is_set(): + state.final_answer = "(interrupted)" + state.add_step(AgentStep(index=index, action="cancelled", tool_name=call.name)) + self.logger.log("cancelled", {"step_index": index, "tool": call.name}, + session_id=session_id, run_id=run_id) + return self._result( + state=state, final_answer="(interrupted)", run_id=run_id, + session_id=session_id, workflow_key=workflow.key if workflow else None, + active_rule_keys=active_rule_keys_for_learning, matched_skills=matched_skills, + ) try: result = self.tools.run(call.name, tool_context, call.arguments) serialized = serialize_tool_result(result) diff --git a/agentic_loop/factory.py b/agentic_loop/factory.py index 044246b..3d2794e 100644 --- a/agentic_loop/factory.py +++ b/agentic_loop/factory.py @@ -1,6 +1,7 @@ import os from pathlib import Path +from .context import ContextBuilder from .controller import AgentController from .logs import CompositeTraceLogger, JsonlTraceLogger, SQLiteTraceLogger from .memory import JsonlMemory @@ -42,6 +43,7 @@ def create_controller( workflow_key: str | None = None, learning_mode: str = "draft", learning_threshold: int = 2, + system_prompt: str | None = None, ) -> AgentController: workspace_path = Path(workspace) registry = tools or create_default_tools(enable_network=enable_network_tools) @@ -108,6 +110,7 @@ def create_controller( return AgentController( model=selected_model, + context_builder=ContextBuilder(system_prompt=system_prompt) if system_prompt else None, tools=registry, policy=WorkspacePolicy( workspace_path, diff --git a/agentic_loop/server.py b/agentic_loop/server.py index 112b306..dfe03ab 100644 --- a/agentic_loop/server.py +++ b/agentic_loop/server.py @@ -75,6 +75,7 @@ def __init__( voice_token_client: GeminiLiveTokenClient | None = None, learning_mode: str = "draft", learning_threshold: int = 2, + system_prompt: str | None = None, ): if voice_provider not in VOICE_PROVIDERS: allowed = ", ".join(sorted(VOICE_PROVIDERS)) @@ -102,6 +103,7 @@ def __init__( self.voice_model = voice_model or DEFAULT_GEMINI_LIVE_MODEL self.learning_mode = learning_mode self.learning_threshold = learning_threshold + self.system_prompt = system_prompt self.gemini_api_key = ( gemini_api_key or os.environ.get("GEMINI_API_KEY") @@ -183,6 +185,7 @@ def _make_controller(self, model_selection: ModelSelection | None = None): tools=self._create_tools(), learning_mode=self.learning_mode, learning_threshold=self.learning_threshold, + system_prompt=self.system_prompt, ) def _make_session( @@ -224,6 +227,7 @@ def chat( disabled_rule_keys: list[str] | None = None, workflow_key: str | None = None, model_selection: ModelSelection | None = None, + cancel_event=None, ) -> dict[str, Any]: session_id, session = self.get_session(session_id, model_selection) selected_model = self.session_model_selections.get(session_id, self.default_model_selection) @@ -232,6 +236,7 @@ def chat( enabled_rule_keys=enabled_rule_keys, disabled_rule_keys=disabled_rule_keys, workflow_key=workflow_key or self.session_workflows.get(session_id), + cancel_event=cancel_event, ) payload = serialize_result(result, selected_model) payload["session_id"] = session_id diff --git a/agentic_loop/session.py b/agentic_loop/session.py index 5aaae3e..603614f 100644 --- a/agentic_loop/session.py +++ b/agentic_loop/session.py @@ -18,6 +18,7 @@ def ask( enabled_rule_keys: list[str] | tuple[str, ...] | set[str] | None = None, disabled_rule_keys: list[str] | tuple[str, ...] | set[str] | None = None, workflow_key: str | None = None, + cancel_event=None, ) -> AgentResult: result = self.controller.run( user_message, @@ -26,6 +27,7 @@ def ask( enabled_rule_keys=enabled_rule_keys, disabled_rule_keys=disabled_rule_keys, workflow_key=workflow_key, + cancel_event=cancel_event, ) user = Message(role="user", content=user_message) assistant = Message(role="assistant", content=result.final_answer) diff --git a/agentic_loop/tools.py b/agentic_loop/tools.py index 70566e5..9b191c9 100644 --- a/agentic_loop/tools.py +++ b/agentic_loop/tools.py @@ -1,6 +1,8 @@ import csv +import ipaddress import json import re +import socket import urllib.parse import urllib.request import xml.etree.ElementTree as ET @@ -12,6 +14,41 @@ from .memory import MemoryStore +class UnsafeUrlError(ValueError): + """A web request was rejected (bad scheme or a non-public/SSRF target).""" + + +def _ip_is_public(ip: str) -> bool: + try: + a = ipaddress.ip_address(ip) + except ValueError: + return False + return not (a.is_loopback or a.is_private or a.is_link_local + or a.is_reserved or a.is_multicast or a.is_unspecified) + + +def require_public_http_url(url: str) -> str: + """SSRF guard: allow only http(s) URLs whose host resolves to PUBLIC IPs. + Blocks loopback (the bundled Ollama / the agent's own API), link-local cloud + metadata (169.254.169.254), and private LAN services.""" + parsed = urllib.parse.urlparse((url or "").strip()) + if parsed.scheme.lower() not in {"http", "https"}: + raise UnsafeUrlError(f"unsupported URL scheme: {parsed.scheme!r}") + host = parsed.hostname or "" + try: + ipaddress.ip_address(host) + ok = _ip_is_public(host) + except ValueError: + try: + infos = socket.getaddrinfo(host, None) + except Exception as exc: # noqa: BLE001 + raise UnsafeUrlError(f"could not resolve host {host!r}: {exc}") from exc + ok = bool(infos) and all(_ip_is_public(i[4][0]) for i in infos) + if not ok: + raise UnsafeUrlError(f"refusing a non-public host: {host!r}") + return url + + class WebClient(Protocol): def get_json(self, url: str, timeout_s: float = 10.0) -> Any: ... @@ -26,6 +63,7 @@ def get_json(self, url: str, timeout_s: float = 10.0) -> Any: return json.loads(text) def get_text(self, url: str, timeout_s: float = 10.0) -> str: + require_public_http_url(url) # SSRF guard before any network call request = urllib.request.Request( url, headers={"User-Agent": "agentic-loop/0.1"},