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.
- agentshim/__init__.py +24 -0
- agentshim/base.py +52 -0
- agentshim/claude/__init__.py +3 -0
- agentshim/claude/agent.py +264 -0
- agentshim/claude/events.py +149 -0
- agentshim/claude/hooks/__init__.py +0 -0
- agentshim/claude/hooks/confine_reads.py +77 -0
- agentshim/claude_events.py +21 -0
- agentshim/cli_agent.py +433 -0
- agentshim/codex/__init__.py +3 -0
- agentshim/codex/agent.py +236 -0
- agentshim/codex/events.py +216 -0
- agentshim/codex_events.py +23 -0
- agentshim/events.py +24 -0
- agentshim/gemini/__init__.py +3 -0
- agentshim/gemini/agent.py +234 -0
- agentshim/gemini/events.py +78 -0
- agentshim/gemini_events.py +11 -0
- agentshim/llm_client.py +64 -0
- agentshim/mcp_config.py +25 -0
- agentshim/opencode/__init__.py +3 -0
- agentshim/opencode/agent.py +196 -0
- agentshim/opencode/events.py +90 -0
- agentshim/opencode_events.py +11 -0
- agentshim/py.typed +0 -0
- agentshim/sandbox.py +133 -0
- agentshim/subagent.py +93 -0
- agentshim/trajectory.py +168 -0
- agentshim/usage.py +63 -0
- agentshim/utils.py +120 -0
- agentshim-0.1.0.dist-info/METADATA +61 -0
- agentshim-0.1.0.dist-info/RECORD +33 -0
- agentshim-0.1.0.dist-info/WHEEL +4 -0
|
@@ -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,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
|
+
]
|
agentshim/llm_client.py
ADDED
|
@@ -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
|
agentshim/mcp_config.py
ADDED
|
@@ -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
|