agentshim 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,216 @@
1
+ from __future__ import annotations
2
+
3
+ from abc import ABC, abstractmethod
4
+ from typing import Any, cast
5
+
6
+ from ..utils import truncate_content, truncate_tool_params
7
+
8
+
9
+ class CodexEvent(ABC):
10
+ """Base class for Codex (``codex exec --json``) stream events."""
11
+
12
+ @abstractmethod
13
+ def render(self, log_prefix: str) -> str | None:
14
+ """Render the event as a string for terminal output."""
15
+
16
+ @staticmethod
17
+ def from_dict(data: dict[str, Any]) -> CodexEvent | None:
18
+ """Factory method to create events from JSON data."""
19
+ event_type = data.get("type")
20
+
21
+ if event_type == "thread.started":
22
+ thread_id_raw = data.get("thread_id")
23
+ thread_id = thread_id_raw if isinstance(thread_id_raw, str) else None
24
+ return ThreadStartedEvent(thread_id=thread_id)
25
+
26
+ if event_type == "turn.started":
27
+ return LifecycleEvent(event_type)
28
+
29
+ if event_type == "turn.completed":
30
+ usage_raw = data.get("usage")
31
+ usage = cast("dict[str, Any]", usage_raw) if isinstance(usage_raw, dict) else {}
32
+ return TurnCompletedEvent(
33
+ input_tokens=int(usage.get("input_tokens") or 0),
34
+ cached_input_tokens=int(usage.get("cached_input_tokens") or 0),
35
+ output_tokens=int(usage.get("output_tokens") or 0),
36
+ )
37
+
38
+ if event_type in ("item.started", "item.completed"):
39
+ item_raw = data.get("item")
40
+ if not isinstance(item_raw, dict):
41
+ return None
42
+ item = cast("dict[str, Any]", item_raw)
43
+ item_type_raw = item.get("type")
44
+ item_type = item_type_raw if isinstance(item_type_raw, str) else None
45
+ item_id_raw = item.get("id")
46
+ item_id = item_id_raw if isinstance(item_id_raw, str) else None
47
+ completed = event_type == "item.completed"
48
+
49
+ if item_type == "agent_message":
50
+ if not completed:
51
+ return None
52
+ text_raw = item.get("text", "")
53
+ return TextEvent(text=text_raw if isinstance(text_raw, str) else "")
54
+
55
+ if item_type == "command_execution":
56
+ exit_code_raw = item.get("exit_code")
57
+ exit_code = exit_code_raw if isinstance(exit_code_raw, int) else None
58
+ status_raw = item.get("status")
59
+ status = status_raw if isinstance(status_raw, str) else None
60
+ if completed:
61
+ return ToolResultEvent(
62
+ tool_id=item_id,
63
+ output=item.get("aggregated_output", ""),
64
+ exit_code=exit_code,
65
+ status=status,
66
+ )
67
+ command_raw = item.get("command", "")
68
+ command = command_raw if isinstance(command_raw, str) else ""
69
+ return ToolUseEvent(
70
+ tool_id=item_id,
71
+ tool_name="shell",
72
+ parameters={"command": command},
73
+ )
74
+
75
+ status_raw = item.get("status")
76
+ status = status_raw if isinstance(status_raw, str) else None
77
+ if completed:
78
+ return ToolResultEvent(
79
+ tool_id=item_id,
80
+ output=_summarize_item(item),
81
+ exit_code=None,
82
+ status=status,
83
+ )
84
+ return ToolUseEvent(
85
+ tool_id=item_id,
86
+ tool_name=item_type or "item",
87
+ parameters=_item_parameters(item),
88
+ )
89
+
90
+ if event_type == "turn.failed":
91
+ err_raw = data.get("error")
92
+ if isinstance(err_raw, dict):
93
+ err = cast("dict[str, Any]", err_raw)
94
+ msg_raw = err.get("message", "")
95
+ message = msg_raw if isinstance(msg_raw, str) else str(msg_raw)
96
+ else:
97
+ message = str(err_raw) if err_raw else ""
98
+ return ErrorEvent(message=message)
99
+
100
+ if event_type == "error":
101
+ message_raw = data.get("message", "")
102
+ return ErrorEvent(message=message_raw if isinstance(message_raw, str) else str(message_raw))
103
+
104
+ return None
105
+
106
+
107
+ def _item_parameters(item: dict[str, Any]) -> dict[str, Any]:
108
+ """Extract a parameter dict from a generic codex item payload."""
109
+ excluded = {"id", "type", "status"}
110
+ return {k: v for k, v in item.items() if k not in excluded}
111
+
112
+
113
+ def _summarize_item(item: dict[str, Any]) -> str:
114
+ """Summarize a generic codex item for the tool-result output field."""
115
+ for key in ("text", "summary", "output", "result"):
116
+ value = item.get(key)
117
+ if isinstance(value, str) and value:
118
+ return value
119
+ return ""
120
+
121
+
122
+ class ThreadStartedEvent(CodexEvent):
123
+ """Thread-level start event carrying the resumable ``thread_id``."""
124
+
125
+ def __init__(self, thread_id: str | None):
126
+ self.thread_id = thread_id
127
+
128
+ def render(self, log_prefix: str) -> str | None:
129
+ return None
130
+
131
+
132
+ class LifecycleEvent(CodexEvent):
133
+ """Turn lifecycle marker; not rendered."""
134
+
135
+ def __init__(self, event_type: str):
136
+ self.event_type = event_type
137
+
138
+ def render(self, log_prefix: str) -> str | None:
139
+ return None
140
+
141
+
142
+ class TextEvent(CodexEvent):
143
+ """Assistant text content event (``agent_message`` item)."""
144
+
145
+ def __init__(self, text: str):
146
+ self.text = text
147
+
148
+ def render(self, log_prefix: str) -> str | None:
149
+ return self.text
150
+
151
+
152
+ class ToolUseEvent(CodexEvent):
153
+ """Tool call start event (``item.started``)."""
154
+
155
+ def __init__(self, tool_name: str, tool_id: str | None, parameters: Any):
156
+ self.tool_name = tool_name
157
+ self.tool_id = tool_id
158
+ self.parameters = parameters
159
+
160
+ def render(self, log_prefix: str) -> str:
161
+ truncated = truncate_tool_params(self.tool_name, self.parameters)
162
+ return f"{log_prefix} \033[34m[Tool Use] {self.tool_name} {truncated}\033[0m"
163
+
164
+
165
+ class ToolResultEvent(CodexEvent):
166
+ """Tool call completion event (``item.completed``)."""
167
+
168
+ def __init__(
169
+ self,
170
+ output: Any,
171
+ tool_id: str | None,
172
+ exit_code: int | None = None,
173
+ status: str | None = None,
174
+ ):
175
+ if isinstance(output, list):
176
+ items = cast("list[Any]", output)
177
+ self.output = "\n".join(str(item) for item in items)
178
+ else:
179
+ self.output = str(output) if output else ""
180
+ self.tool_id = tool_id
181
+ self.exit_code = exit_code
182
+ self.status = status
183
+ self.tool_name_resolved: str = "Tool"
184
+
185
+ def render(self, log_prefix: str) -> str:
186
+ if not self.output:
187
+ return f"{log_prefix} \033[32m{self.tool_name_resolved} ran successfully\033[0m"
188
+ truncated = truncate_content(self.output)
189
+ return f"{log_prefix} \033[32m[Tool Result] {truncated}\033[0m"
190
+
191
+
192
+ class TurnCompletedEvent(CodexEvent):
193
+ """Per-turn usage summary emitted by Codex."""
194
+
195
+ def __init__(
196
+ self,
197
+ input_tokens: int = 0,
198
+ cached_input_tokens: int = 0,
199
+ output_tokens: int = 0,
200
+ ):
201
+ self.input_tokens = input_tokens
202
+ self.cached_input_tokens = cached_input_tokens
203
+ self.output_tokens = output_tokens
204
+
205
+ def render(self, log_prefix: str) -> str | None:
206
+ return None
207
+
208
+
209
+ class ErrorEvent(CodexEvent):
210
+ """Error event emitted on turn failure or top-level error."""
211
+
212
+ def __init__(self, message: str):
213
+ self.message = message
214
+
215
+ def render(self, log_prefix: str) -> str:
216
+ return f"{log_prefix} \033[31m[Error] {self.message}\033[0m"
@@ -0,0 +1,23 @@
1
+ """Compatibility shim for older ``agentshim.codex_events`` imports."""
2
+
3
+ from .codex.events import (
4
+ CodexEvent,
5
+ ErrorEvent,
6
+ LifecycleEvent,
7
+ TextEvent,
8
+ ThreadStartedEvent,
9
+ ToolResultEvent,
10
+ ToolUseEvent,
11
+ TurnCompletedEvent,
12
+ )
13
+
14
+ __all__ = [
15
+ "CodexEvent",
16
+ "ErrorEvent",
17
+ "LifecycleEvent",
18
+ "TextEvent",
19
+ "ThreadStartedEvent",
20
+ "ToolResultEvent",
21
+ "ToolUseEvent",
22
+ "TurnCompletedEvent",
23
+ ]
agentshim/events.py ADDED
@@ -0,0 +1,24 @@
1
+ from typing import Any, Protocol
2
+
3
+
4
+ class AgentEventHandler(Protocol):
5
+ """Protocol for handling agent events."""
6
+
7
+ def on_thinking(self, text: str) -> None:
8
+ """Handle agent thinking output."""
9
+ ...
10
+
11
+ def on_tool_call(self, tool: str, args: dict[str, Any] | str | None = None) -> None:
12
+ """Handle tool execution start."""
13
+ ...
14
+
15
+ def on_tool_result(
16
+ self,
17
+ tool: str,
18
+ stdout: str = "",
19
+ stderr: str = "",
20
+ exit_code: int | None = None,
21
+ duration: float | None = None,
22
+ ) -> None:
23
+ """Handle tool execution result."""
24
+ ...
@@ -0,0 +1,3 @@
1
+ from .agent import GeminiCodingAgent, GeminiGenerationSession
2
+
3
+ __all__ = ["GeminiCodingAgent", "GeminiGenerationSession"]
@@ -0,0 +1,234 @@
1
+ import json
2
+ import logging
3
+ import subprocess
4
+ import time
5
+ from collections.abc import Callable
6
+ from pathlib import Path
7
+ from typing import Any
8
+
9
+ import agentshim.trajectory as _trajectory_module
10
+ from agentshim.trajectory import TrajectoryRecorderProtocol
11
+
12
+ from ..base import register_provider
13
+ from ..cli_agent import CLICodingAgent, CLIGenerationSession
14
+ from ..events import AgentEventHandler
15
+ from ..sandbox import SandboxConfig
16
+ from ..usage import ProviderUsage, TokenUsage
17
+ from .events import GeminiEvent, InitEvent, MessageEvent, ToolResultEvent, ToolUseEvent
18
+
19
+ _logger = logging.getLogger(__name__)
20
+
21
+
22
+ class GeminiGenerationSession(CLIGenerationSession):
23
+ def __init__(self, **kwargs: Any):
24
+ super().__init__(**kwargs)
25
+ # Initialize state required for stream processing
26
+ self.tool_map: dict[str, str] = {}
27
+ self.tool_start_times: dict[str, float] = {}
28
+ self.tool_args: dict[str, Any] = {}
29
+ # Capture call_id and run_id for correlation
30
+ self.call_id = _trajectory_module.get_current_call_id()
31
+ self.run_id = _trajectory_module.get_run_id()
32
+ # Gemini's stream-json does not emit token usage. We count
33
+ # assistant MessageEvents as a turn proxy; token fields stay 0.
34
+ self._assistant_message_count: int = 0
35
+
36
+ def _write_call_metadata(self):
37
+ """Write metadata file to help correlate Gemini session with trajectory call."""
38
+ if self.call_id is None or self.run_id is None:
39
+ return
40
+
41
+ try:
42
+ gemini_tmp_dir = Path.home() / ".gemini" / "tmp"
43
+ if not gemini_tmp_dir.exists():
44
+ return
45
+
46
+ if self.cwd:
47
+ project_name = Path(self.cwd).name
48
+ project_dir = gemini_tmp_dir / project_name / "chats"
49
+ if project_dir.exists():
50
+ metadata_file = project_dir / f"sds_call_{self.call_id:03d}.json"
51
+ metadata = {
52
+ "call_id": self.call_id,
53
+ "run_id": self.run_id,
54
+ "start_time": time.strftime("%Y-%m-%d %H:%M:%S"),
55
+ "cwd": self.cwd,
56
+ }
57
+ with open(metadata_file, "w") as f:
58
+ json.dump(metadata, f, indent=2)
59
+ except Exception:
60
+ _logger.debug("Failed to write Gemini call metadata", exc_info=True)
61
+
62
+ def run(self, prompt: str) -> str:
63
+ """Execute the generation process, writing call metadata first."""
64
+ self._write_call_metadata()
65
+
66
+ if not self.silent:
67
+ self._log_raw(f"{self.log_prefix} Input Prompt:\n{prompt}\n")
68
+
69
+ return super().run(prompt)
70
+
71
+ def _process_stdout(self, line: str) -> None:
72
+ """Process a line from stdout."""
73
+ if not line:
74
+ return
75
+ try:
76
+ data = json.loads(line)
77
+ event = GeminiEvent.from_dict(data)
78
+ if event:
79
+ self._handle_event(event)
80
+ except json.JSONDecodeError:
81
+ if not self.silent:
82
+ if self._at_line_start:
83
+ self._log_raw(f"{self.log_prefix} ")
84
+ self._log_raw(line.rstrip() + "\n")
85
+ self._at_line_start = True
86
+
87
+ def _handle_event(self, event: GeminiEvent):
88
+ """Handle a single parsed Gemini event."""
89
+ self._update_state(event)
90
+
91
+ if not self.silent:
92
+ self._render_event(event)
93
+
94
+ def _update_state(self, event: GeminiEvent):
95
+ """Update internal state based on the event."""
96
+ if isinstance(event, InitEvent):
97
+ if self.session_id is None and event.session_id:
98
+ self.session_id = event.session_id
99
+ return
100
+
101
+ if isinstance(event, MessageEvent):
102
+ if event.role == "assistant":
103
+ self.stdout_lines.append(event.content)
104
+ self._assistant_message_count += 1
105
+ self.usage = ProviderUsage(
106
+ tokens=TokenUsage(turns=self._assistant_message_count),
107
+ provider="gemini",
108
+ )
109
+ if self.event_handler:
110
+ self.event_handler.on_thinking(event.content)
111
+
112
+ elif isinstance(event, ToolUseEvent):
113
+ if event.tool_id:
114
+ self.tool_map[event.tool_id] = event.tool_name
115
+ self.tool_start_times[event.tool_id] = time.time()
116
+ self.tool_args[event.tool_id] = event.parameters
117
+ if self.event_handler:
118
+ self.event_handler.on_tool_call(event.tool_name, event.parameters)
119
+
120
+ elif isinstance(event, ToolResultEvent) and event.tool_id:
121
+ event.tool_name_resolved = self.tool_map.get(event.tool_id, "Tool")
122
+
123
+ start_time = self.tool_start_times.get(event.tool_id)
124
+ duration = time.time() - start_time if start_time else None
125
+ args = self.tool_args.get(event.tool_id, {})
126
+
127
+ self.recorder.add_tool_call(
128
+ tool=event.tool_name_resolved,
129
+ args=args,
130
+ stdout=event.output,
131
+ duration=duration,
132
+ )
133
+
134
+ if self.event_handler:
135
+ self.event_handler.on_tool_result(
136
+ tool=event.tool_name_resolved,
137
+ stdout=event.output,
138
+ duration=duration,
139
+ )
140
+
141
+ def _render_event(self, event: GeminiEvent):
142
+ """Render the event to stdout."""
143
+ if isinstance(event, MessageEvent):
144
+ if event.role == "assistant":
145
+ self._print_stream_content(event.content)
146
+ return
147
+
148
+ if not self._at_line_start:
149
+ self._log_raw("\n")
150
+ self._at_line_start = True
151
+
152
+ output = event.render(self.log_prefix)
153
+ if output:
154
+ self._log_raw(output + "\n")
155
+
156
+
157
+ @register_provider("gemini")
158
+ class GeminiCodingAgent(CLICodingAgent):
159
+ """Coding agent implementation using the Gemini CLI tool."""
160
+
161
+ def __init__(
162
+ self,
163
+ model: str | None = None,
164
+ recorder: TrajectoryRecorderProtocol | None = None,
165
+ event_handler: AgentEventHandler | None = None,
166
+ mcp_servers: list[object] | None = None,
167
+ sandbox: bool | SandboxConfig = False,
168
+ ):
169
+ """Initialize the Gemini coding agent.
170
+
171
+ Args:
172
+ model: Optional model name to use.
173
+ recorder: Trajectory recorder instance.
174
+ event_handler: Optional event handler for UI updates.
175
+ mcp_servers: Optional list of MCP server configurations.
176
+ sandbox: Not supported for Gemini; must be False.
177
+
178
+ Raises:
179
+ ValueError: If mcp_servers is non-empty (not supported).
180
+ NotImplementedError: If ``sandbox`` is truthy.
181
+ """
182
+ if mcp_servers:
183
+ raise ValueError("GeminiCodingAgent does not support programmatic MCP server configuration via CLI flags")
184
+ if sandbox:
185
+ raise NotImplementedError("sandbox is not supported for GeminiCodingAgent")
186
+ super().__init__("gemini", model, recorder, event_handler)
187
+
188
+ @property
189
+ def gemini_path(self) -> str:
190
+ """Return path to gemini binary (for backward compatibility)."""
191
+ return self.binary_path
192
+
193
+ @property
194
+ def _log_prefix(self) -> str:
195
+ """Return the log prefix for this agent."""
196
+ return "[Gemini]"
197
+
198
+ def _get_command(self, prompt: str, resume_session_id: str | None = None) -> list[str]:
199
+ cmd = [self.binary_path]
200
+
201
+ cmd.extend(["-y"])
202
+
203
+ if self.model:
204
+ cmd.extend(["--model", self.model])
205
+
206
+ cmd.extend(["-o", "stream-json"])
207
+
208
+ if resume_session_id:
209
+ cmd.extend(["--resume", resume_session_id])
210
+
211
+ return cmd
212
+
213
+ def _create_session(
214
+ self,
215
+ cmd: list[str],
216
+ cwd: str | None = None,
217
+ timeout: int = 300,
218
+ silent: bool = False,
219
+ recorder: TrajectoryRecorderProtocol | None = None,
220
+ on_process_started: Callable[[subprocess.Popen[str]], None] | None = None,
221
+ ) -> GeminiGenerationSession:
222
+ return GeminiGenerationSession(
223
+ binary_name=self.binary_name,
224
+ env=self.env,
225
+ log_prefix=self._log_prefix,
226
+ cmd=cmd,
227
+ logger=self.logger,
228
+ cwd=cwd,
229
+ timeout=timeout,
230
+ silent=silent,
231
+ recorder=recorder,
232
+ event_handler=self.event_handler,
233
+ on_process_started=on_process_started,
234
+ )
@@ -0,0 +1,78 @@
1
+ from __future__ import annotations
2
+
3
+ from abc import ABC, abstractmethod
4
+ from typing import Any
5
+
6
+ from ..utils import truncate_content, truncate_tool_params
7
+
8
+
9
+ class GeminiEvent(ABC):
10
+ """Base class for Gemini stream events."""
11
+
12
+ @abstractmethod
13
+ def render(self, log_prefix: str) -> str | None:
14
+ """Render the event as a string for terminal output."""
15
+
16
+ @staticmethod
17
+ def from_dict(data: dict[str, Any]) -> GeminiEvent | None:
18
+ """Factory method to create events from JSON data."""
19
+ msg_type = data.get("type")
20
+
21
+ if msg_type == "init":
22
+ return InitEvent(session_id=data.get("session_id"))
23
+ if msg_type == "message":
24
+ return MessageEvent(role=data.get("role", ""), content=data.get("content", ""))
25
+ if msg_type == "tool_use":
26
+ return ToolUseEvent(
27
+ tool_name=data.get("tool_name", "Tool"),
28
+ tool_id=data.get("tool_id"),
29
+ parameters=data.get("parameters"),
30
+ )
31
+ if msg_type == "tool_result":
32
+ return ToolResultEvent(output=data.get("output", ""), tool_id=data.get("tool_id"))
33
+ return None
34
+
35
+
36
+ class InitEvent(GeminiEvent):
37
+ """Session initialization event; carries the resumable ``session_id``."""
38
+
39
+ def __init__(self, session_id: str | None):
40
+ self.session_id = session_id
41
+
42
+ def render(self, log_prefix: str) -> str | None:
43
+ return None
44
+
45
+
46
+ class MessageEvent(GeminiEvent):
47
+ def __init__(self, role: str, content: str):
48
+ self.role = role
49
+ self.content = content
50
+
51
+ def render(self, log_prefix: str) -> str | None:
52
+ # Message rendering is handled specially due to streaming
53
+ # This is just a placeholder or could handle non-streaming blocks
54
+ return self.content
55
+
56
+
57
+ class ToolUseEvent(GeminiEvent):
58
+ def __init__(self, tool_name: str, tool_id: str | None, parameters: Any):
59
+ self.tool_name = tool_name
60
+ self.tool_id = tool_id
61
+ self.parameters = parameters
62
+
63
+ def render(self, log_prefix: str) -> str:
64
+ truncated = truncate_tool_params(self.tool_name, self.parameters)
65
+ return f"{log_prefix} \033[34m[Tool Use] {self.tool_name} {truncated}\033[0m"
66
+
67
+
68
+ class ToolResultEvent(GeminiEvent):
69
+ def __init__(self, output: str, tool_id: str | None):
70
+ self.output = output
71
+ self.tool_id = tool_id
72
+ self.tool_name_resolved: str = "Tool" # To be set externally
73
+
74
+ def render(self, log_prefix: str) -> str:
75
+ if not self.output:
76
+ return f"{log_prefix} \033[32m{self.tool_name_resolved} ran successfully\033[0m"
77
+ truncated = truncate_content(self.output)
78
+ return f"{log_prefix} \033[32m[Tool Result] {truncated}\033[0m"
@@ -0,0 +1,11 @@
1
+ """Compatibility shim for older ``agentshim.gemini_events`` imports."""
2
+
3
+ from .gemini.events import GeminiEvent, InitEvent, MessageEvent, ToolResultEvent, ToolUseEvent
4
+
5
+ __all__ = [
6
+ "GeminiEvent",
7
+ "InitEvent",
8
+ "MessageEvent",
9
+ "ToolResultEvent",
10
+ "ToolUseEvent",
11
+ ]
@@ -0,0 +1,64 @@
1
+ """Thin litellm wrapper with automatic token tracking and trajectory recording."""
2
+
3
+ import os
4
+ from typing import Any
5
+
6
+ from agentshim.subagent import litellm_call_with_retry
7
+ from agentshim.trajectory import TrajectoryRecorderProtocol
8
+
9
+
10
+ class LiteLLMClient:
11
+ """litellm completion wrapper that accumulates token usage and records to trajectory.
12
+
13
+ Every ``complete()`` call:
14
+ - Delegates to ``litellm_call_with_retry`` for retry-on-network-error logic.
15
+ - Accumulates prompt/completion/total tokens in ``_token_usage``.
16
+ - Calls ``recorder.record_token_usage()`` with the running total.
17
+
18
+ The caller is responsible for managing conversation history (``messages``).
19
+ This avoids coupling: pass the full list each call, append results yourself.
20
+ """
21
+
22
+ def __init__(
23
+ self,
24
+ model: str,
25
+ location: str | None = None,
26
+ recorder: TrajectoryRecorderProtocol | None = None,
27
+ ):
28
+ self.model = model
29
+ self.location = location
30
+ self.recorder = recorder
31
+ self._token_usage: dict[str, int] = {
32
+ "prompt_tokens": 0,
33
+ "completion_tokens": 0,
34
+ "total_tokens": 0,
35
+ }
36
+
37
+ def complete(self, messages: list[dict[str, Any]], label: str = "llm call") -> str:
38
+ """Call litellm with *messages*, accumulate tokens, and record to trajectory.
39
+
40
+ Args:
41
+ messages: Full conversation history to send (caller-managed).
42
+ label: Human-readable label for retry/logging messages.
43
+
44
+ Returns:
45
+ Assistant response text.
46
+
47
+ Raises:
48
+ Exception: Any non-retried litellm error propagates to the caller.
49
+ """
50
+ kwargs: dict[str, Any] = {
51
+ "model": self.model,
52
+ "messages": messages,
53
+ "cache": {"no-cache": True},
54
+ }
55
+ loc = self.location or os.environ.get("VERTEX_LOCATION")
56
+ if loc:
57
+ kwargs["vertex_location"] = loc
58
+
59
+ result = litellm_call_with_retry(kwargs, label=label, token_acc=self._token_usage)
60
+
61
+ if self.recorder and hasattr(self.recorder, "record_token_usage"):
62
+ self.recorder.record_token_usage(self._token_usage.copy()) # type: ignore[reportArgumentType]
63
+
64
+ return result
@@ -0,0 +1,25 @@
1
+ from pydantic import BaseModel, ConfigDict, Field
2
+
3
+
4
+ class HttpMcpServer(BaseModel):
5
+ """MCP server accessed over HTTP/SSE."""
6
+
7
+ model_config = ConfigDict(frozen=True, extra="forbid")
8
+
9
+ name: str = Field(min_length=1)
10
+ url: str = Field(min_length=1)
11
+ headers: dict[str, str] = Field(default_factory=dict)
12
+
13
+
14
+ class StdioMcpServer(BaseModel):
15
+ """MCP server launched as a subprocess (stdio transport)."""
16
+
17
+ model_config = ConfigDict(frozen=True, extra="forbid")
18
+
19
+ name: str = Field(min_length=1)
20
+ command: str = Field(min_length=1)
21
+ args: list[str] = Field(default_factory=list)
22
+ env: dict[str, str] = Field(default_factory=dict)
23
+
24
+
25
+ McpServerConfig = HttpMcpServer | StdioMcpServer
@@ -0,0 +1,3 @@
1
+ from .agent import OpencodeCodingAgent, OpencodeGenerationSession
2
+
3
+ __all__ = ["OpencodeCodingAgent", "OpencodeGenerationSession"]