sky-dev 0.0.5__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.
- sky/__init__.py +6 -0
- sky/cli.py +833 -0
- sky/config/__init__.py +19 -0
- sky/config/models.yaml +39 -0
- sky/config/schema.py +275 -0
- sky/core/__init__.py +48 -0
- sky/core/approval.py +92 -0
- sky/core/benchmark.py +39 -0
- sky/core/chat.py +211 -0
- sky/core/fast_loop.py +330 -0
- sky/core/mode_prompts.py +106 -0
- sky/core/router.py +264 -0
- sky/core/subagent.py +302 -0
- sky/core/workflow.py +662 -0
- sky/errors.py +29 -0
- sky/memory/__init__.py +27 -0
- sky/memory/indexer.py +450 -0
- sky/memory/vectorstore.py +271 -0
- sky/security/__init__.py +17 -0
- sky/security/audit.py +23 -0
- sky/security/detection.py +94 -0
- sky/security/guardrails.py +106 -0
- sky/security/prompts.py +58 -0
- sky/security/rate_limit.py +36 -0
- sky/security/sanitize.py +152 -0
- sky/storage/__init__.py +21 -0
- sky/storage/db.py +492 -0
- sky/tools/__init__.py +26 -0
- sky/tools/fs_tools.py +139 -0
- sky/tools/git_tools.py +139 -0
- sky/tools/registry.py +219 -0
- sky/tools/search_tools.py +128 -0
- sky/tools/shell_tools.py +73 -0
- sky_dev-0.0.5.dist-info/METADATA +119 -0
- sky_dev-0.0.5.dist-info/RECORD +39 -0
- sky_dev-0.0.5.dist-info/WHEEL +5 -0
- sky_dev-0.0.5.dist-info/entry_points.txt +2 -0
- sky_dev-0.0.5.dist-info/licenses/LICENSE +21 -0
- sky_dev-0.0.5.dist-info/top_level.txt +1 -0
sky/core/router.py
ADDED
|
@@ -0,0 +1,264 @@
|
|
|
1
|
+
"""Model Router with F11 Batching Verification."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Tuple
|
|
4
|
+
|
|
5
|
+
import httpx
|
|
6
|
+
from openai import AsyncOpenAI
|
|
7
|
+
|
|
8
|
+
from sky.config import ModelAssignmentConfig, ModelRoutingConfig
|
|
9
|
+
from sky.storage import DatabaseManager
|
|
10
|
+
from sky.errors import ProviderError, ConfigurationError
|
|
11
|
+
|
|
12
|
+
class ModelRouter:
|
|
13
|
+
"""Routes requests to appropriate LLMs with F11 verification."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, config: ModelRoutingConfig, db: DatabaseManager, session_id: str):
|
|
16
|
+
import os
|
|
17
|
+
self.config = config
|
|
18
|
+
self.db = db
|
|
19
|
+
self.session_id = session_id
|
|
20
|
+
|
|
21
|
+
self.clients = {}
|
|
22
|
+
self.provider_keys = {}
|
|
23
|
+
self.current_key_idx = {}
|
|
24
|
+
|
|
25
|
+
for provider_name, provider_config in self.config.providers.items():
|
|
26
|
+
if provider_config.requires_api_key:
|
|
27
|
+
# Find keys for this provider, e.g., GROQ_API_KEY, NIM_API_KEY
|
|
28
|
+
env_prefix = provider_name.upper()
|
|
29
|
+
keys = []
|
|
30
|
+
for key_name in [f"{env_prefix}_API_KEY", f"{env_prefix}_API_KEY_2", f"{env_prefix}_API_KEY_3"]:
|
|
31
|
+
key = os.getenv(key_name)
|
|
32
|
+
if key:
|
|
33
|
+
keys.append(key)
|
|
34
|
+
if not keys:
|
|
35
|
+
# Log a warning, but don't crash unless they try to use it
|
|
36
|
+
pass
|
|
37
|
+
else:
|
|
38
|
+
self.provider_keys[provider_name] = keys
|
|
39
|
+
self.current_key_idx[provider_name] = 0
|
|
40
|
+
else:
|
|
41
|
+
self.provider_keys[provider_name] = ["dummy"]
|
|
42
|
+
self.current_key_idx[provider_name] = 0
|
|
43
|
+
|
|
44
|
+
self._init_client(provider_name)
|
|
45
|
+
|
|
46
|
+
def _init_client(self, provider_name: str):
|
|
47
|
+
"""Initialize the client for a specific provider."""
|
|
48
|
+
provider_config = self.config.providers.get(provider_name)
|
|
49
|
+
if not provider_config:
|
|
50
|
+
return
|
|
51
|
+
|
|
52
|
+
keys = self.provider_keys.get(provider_name, [])
|
|
53
|
+
if not keys:
|
|
54
|
+
return
|
|
55
|
+
|
|
56
|
+
idx = self.current_key_idx.get(provider_name, 0)
|
|
57
|
+
api_key = keys[idx]
|
|
58
|
+
|
|
59
|
+
kwargs = {
|
|
60
|
+
"api_key": api_key,
|
|
61
|
+
"timeout": httpx.Timeout(provider_config.timeout)
|
|
62
|
+
}
|
|
63
|
+
if provider_config.base_url:
|
|
64
|
+
kwargs["base_url"] = provider_config.base_url
|
|
65
|
+
|
|
66
|
+
self.clients[provider_name] = AsyncOpenAI(**kwargs)
|
|
67
|
+
|
|
68
|
+
def get_model_for_role(self, role: str) -> ModelAssignmentConfig:
|
|
69
|
+
"""Get model config for a specific role."""
|
|
70
|
+
if role in self.config.roles:
|
|
71
|
+
return self.config.roles[role]
|
|
72
|
+
return self.get_fallback_model()
|
|
73
|
+
|
|
74
|
+
def get_fallback_model(self) -> ModelAssignmentConfig:
|
|
75
|
+
"""Get the fallback model config."""
|
|
76
|
+
if self.config.fallback:
|
|
77
|
+
return self.config.fallback
|
|
78
|
+
return ModelAssignmentConfig(provider="nim", model_id="meta/llama-3.1-8b-instruct", reasoning_effort="none")
|
|
79
|
+
|
|
80
|
+
def build_completion_params(
|
|
81
|
+
self,
|
|
82
|
+
assignment: ModelAssignmentConfig,
|
|
83
|
+
messages: List[Dict[str, Any]],
|
|
84
|
+
tools: Optional[List[Dict[str, Any]]],
|
|
85
|
+
tool_choice: str,
|
|
86
|
+
parallel_tool_calls: bool
|
|
87
|
+
) -> Dict[str, Any]:
|
|
88
|
+
"""Build parameters for Groq API call."""
|
|
89
|
+
params: Dict[str, Any] = {
|
|
90
|
+
"model": assignment.model_id,
|
|
91
|
+
"messages": messages,
|
|
92
|
+
}
|
|
93
|
+
if assignment.temperature is not None:
|
|
94
|
+
params["temperature"] = assignment.temperature
|
|
95
|
+
if assignment.max_tokens is not None:
|
|
96
|
+
params["max_tokens"] = assignment.max_tokens
|
|
97
|
+
|
|
98
|
+
if tools:
|
|
99
|
+
params["tools"] = tools
|
|
100
|
+
params["tool_choice"] = tool_choice
|
|
101
|
+
params["parallel_tool_calls"] = parallel_tool_calls
|
|
102
|
+
|
|
103
|
+
return params
|
|
104
|
+
|
|
105
|
+
async def _call_with_rotation(self, params: Dict[str, Any], provider_name: str) -> Any:
|
|
106
|
+
"""Call API, rotating keys if rate limit is hit."""
|
|
107
|
+
from rich.console import Console
|
|
108
|
+
|
|
109
|
+
client = self.clients.get(provider_name)
|
|
110
|
+
if not client:
|
|
111
|
+
raise ConfigurationError(
|
|
112
|
+
f"No active client for provider: {provider_name}.",
|
|
113
|
+
code="SKY-001",
|
|
114
|
+
suggestion=f"Add {provider_name.upper()}_API_KEY to .env or run `sky init`"
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
keys = self.provider_keys.get(provider_name, [])
|
|
118
|
+
last_error = None
|
|
119
|
+
|
|
120
|
+
for _ in range(max(1, len(keys))):
|
|
121
|
+
try:
|
|
122
|
+
return await client.chat.completions.create(**params)
|
|
123
|
+
except Exception as e:
|
|
124
|
+
error_str = str(e).lower()
|
|
125
|
+
if any(err in error_str for err in ["413", "429", "rate limit", "rate_limit_exceeded"]):
|
|
126
|
+
last_error = e
|
|
127
|
+
if len(keys) > 1:
|
|
128
|
+
next_idx = (self.current_key_idx[provider_name] + 1) % len(keys)
|
|
129
|
+
Console().print(f"[yellow]{provider_name.upper()} API key rate limited. Rotating to next key (slot {next_idx + 1}/{len(keys)})...[/yellow]")
|
|
130
|
+
self.current_key_idx[provider_name] = next_idx
|
|
131
|
+
self._init_client(provider_name)
|
|
132
|
+
client = self.clients[provider_name]
|
|
133
|
+
else:
|
|
134
|
+
break
|
|
135
|
+
else:
|
|
136
|
+
raise ProviderError(
|
|
137
|
+
f"Provider {provider_name} request failed: {e}",
|
|
138
|
+
code="SKY-002",
|
|
139
|
+
suggestion="Check your internet connection or run `sky check-providers` to verify availability."
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
raise ProviderError(
|
|
143
|
+
f"Rate limit exceeded after exhausting all {len(keys)} keys.",
|
|
144
|
+
code="SKY-003",
|
|
145
|
+
suggestion="Wait 60 seconds. Consider using a different model or provider."
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
async def route(
|
|
149
|
+
self,
|
|
150
|
+
role: str,
|
|
151
|
+
messages: List[Dict[str, Any]],
|
|
152
|
+
tools: Optional[List[Dict[str, Any]]] = None,
|
|
153
|
+
expected_tool_calls: Optional[int] = None,
|
|
154
|
+
tool_choice: str = "auto"
|
|
155
|
+
) -> Tuple[Dict[str, Any], bool, str]:
|
|
156
|
+
"""Route request to model, returning (response_message_dict, was_fallback, model_used)."""
|
|
157
|
+
if expected_tool_calls and expected_tool_calls > 1 and tools:
|
|
158
|
+
resp, was_fallback = await self.verify_and_dispatch(role, messages, tools, expected_tool_calls, tool_choice)
|
|
159
|
+
model_used = self.get_fallback_model().model_id if was_fallback else self.get_model_for_role(role).model_id
|
|
160
|
+
return resp, was_fallback, model_used
|
|
161
|
+
|
|
162
|
+
assignment = self.get_model_for_role(role)
|
|
163
|
+
params = self.build_completion_params(assignment, messages, tools, tool_choice, self.config.parallel_tool_calls)
|
|
164
|
+
|
|
165
|
+
chat_completion = await self._call_with_rotation(params, assignment.provider)
|
|
166
|
+
resp = chat_completion.choices[0].message
|
|
167
|
+
|
|
168
|
+
if chat_completion.usage:
|
|
169
|
+
cost = (chat_completion.usage.prompt_tokens + chat_completion.usage.completion_tokens) * 0.000001
|
|
170
|
+
self.db.log_usage(
|
|
171
|
+
self.session_id,
|
|
172
|
+
assignment.model_id,
|
|
173
|
+
role,
|
|
174
|
+
chat_completion.usage.prompt_tokens,
|
|
175
|
+
chat_completion.usage.completion_tokens,
|
|
176
|
+
cost,
|
|
177
|
+
0.0
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
resp_dict: Dict[str, Any] = {
|
|
181
|
+
"role": "assistant",
|
|
182
|
+
"content": resp.content,
|
|
183
|
+
}
|
|
184
|
+
if resp.tool_calls:
|
|
185
|
+
resp_dict["tool_calls"] = [
|
|
186
|
+
{
|
|
187
|
+
"id": tc.id,
|
|
188
|
+
"type": "function",
|
|
189
|
+
"function": {
|
|
190
|
+
"name": tc.function.name,
|
|
191
|
+
"arguments": tc.function.arguments
|
|
192
|
+
}
|
|
193
|
+
} for tc in resp.tool_calls
|
|
194
|
+
]
|
|
195
|
+
return resp_dict, False, assignment.model_id
|
|
196
|
+
|
|
197
|
+
async def verify_and_dispatch(
|
|
198
|
+
self,
|
|
199
|
+
role: str,
|
|
200
|
+
messages: List[Dict[str, Any]],
|
|
201
|
+
tools: List[Dict[str, Any]],
|
|
202
|
+
expected_tool_calls: int = 1,
|
|
203
|
+
tool_choice: str = "auto"
|
|
204
|
+
) -> Tuple[Dict[str, Any], bool]:
|
|
205
|
+
"""F11 PRIMARY FUNCTION: Verify tool batching count and fallback if under-calling."""
|
|
206
|
+
assignment = self.get_model_for_role(role)
|
|
207
|
+
params = self.build_completion_params(assignment, messages, tools, tool_choice, True)
|
|
208
|
+
|
|
209
|
+
chat_completion = await self._call_with_rotation(params, assignment.provider)
|
|
210
|
+
resp = chat_completion.choices[0].message
|
|
211
|
+
|
|
212
|
+
tool_call_count = len(resp.tool_calls) if resp.tool_calls else 0
|
|
213
|
+
|
|
214
|
+
if tool_call_count < expected_tool_calls:
|
|
215
|
+
self._log_fallback_event(assignment.model_id, tool_call_count, expected_tool_calls, "under_calling")
|
|
216
|
+
|
|
217
|
+
fallback = self.get_fallback_model()
|
|
218
|
+
fb_params = self.build_completion_params(fallback, messages, tools, tool_choice, True)
|
|
219
|
+
fb_completion = await self._call_with_rotation(fb_params, fallback.provider)
|
|
220
|
+
fb_resp = fb_completion.choices[0].message
|
|
221
|
+
|
|
222
|
+
resp_dict: Dict[str, Any] = {
|
|
223
|
+
"role": "assistant",
|
|
224
|
+
"content": fb_resp.content,
|
|
225
|
+
}
|
|
226
|
+
if fb_resp.tool_calls:
|
|
227
|
+
resp_dict["tool_calls"] = [
|
|
228
|
+
{
|
|
229
|
+
"id": tc.id,
|
|
230
|
+
"type": "function",
|
|
231
|
+
"function": {
|
|
232
|
+
"name": tc.function.name,
|
|
233
|
+
"arguments": tc.function.arguments
|
|
234
|
+
}
|
|
235
|
+
} for tc in fb_resp.tool_calls
|
|
236
|
+
]
|
|
237
|
+
return resp_dict, True
|
|
238
|
+
|
|
239
|
+
resp_dict = {
|
|
240
|
+
"role": "assistant",
|
|
241
|
+
"content": resp.content,
|
|
242
|
+
}
|
|
243
|
+
if resp.tool_calls:
|
|
244
|
+
resp_dict["tool_calls"] = [
|
|
245
|
+
{
|
|
246
|
+
"id": tc.id,
|
|
247
|
+
"type": "function",
|
|
248
|
+
"function": {
|
|
249
|
+
"name": tc.function.name,
|
|
250
|
+
"arguments": tc.function.arguments
|
|
251
|
+
}
|
|
252
|
+
} for tc in resp.tool_calls
|
|
253
|
+
]
|
|
254
|
+
return resp_dict, False
|
|
255
|
+
|
|
256
|
+
def _log_fallback_event(self, model: str, actual: int, expected: int, reason: str) -> None:
|
|
257
|
+
"""Log fallback event to audit log."""
|
|
258
|
+
self.db._write_audit_log(self.session_id, {
|
|
259
|
+
"event": "fallback_triggered",
|
|
260
|
+
"model": model,
|
|
261
|
+
"actual_tool_calls": actual,
|
|
262
|
+
"expected_tool_calls": expected,
|
|
263
|
+
"reason": reason
|
|
264
|
+
})
|
sky/core/subagent.py
ADDED
|
@@ -0,0 +1,302 @@
|
|
|
1
|
+
"""Subagent system for SKY workflow engine."""
|
|
2
|
+
|
|
3
|
+
from enum import Enum
|
|
4
|
+
from typing import Dict, List, Any, Optional, Literal
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
import logging
|
|
7
|
+
import json
|
|
8
|
+
|
|
9
|
+
from sky.config.schema import DexProjectConfig
|
|
10
|
+
from sky.storage.db import DatabaseManager
|
|
11
|
+
from sky.core.router import ModelRouter
|
|
12
|
+
from sky.core.approval import ApprovalGate
|
|
13
|
+
from sky.core.fast_loop import FastLoopEngine
|
|
14
|
+
from sky.memory.indexer import RepoIndexer
|
|
15
|
+
|
|
16
|
+
logger = logging.getLogger(__name__)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class SubagentRole(str, Enum):
|
|
20
|
+
"""Roles for subagents."""
|
|
21
|
+
PLANNER = "planner"
|
|
22
|
+
CODER = "coder"
|
|
23
|
+
TESTER = "tester"
|
|
24
|
+
REVIEWER = "reviewer"
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass
|
|
28
|
+
class SubagentConfig:
|
|
29
|
+
"""Configuration for a subagent."""
|
|
30
|
+
role: SubagentRole
|
|
31
|
+
model_role: str # Corresponds to models.yaml role
|
|
32
|
+
tools: List[str] = field(default_factory=list)
|
|
33
|
+
system_prompt_template: str = ""
|
|
34
|
+
max_turns: int = 10
|
|
35
|
+
temperature: float = 0.3
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
# Role-specific configurations
|
|
39
|
+
ROLE_CONFIGS = {
|
|
40
|
+
SubagentRole.PLANNER: SubagentConfig(
|
|
41
|
+
role=SubagentRole.PLANNER,
|
|
42
|
+
model_role="planning",
|
|
43
|
+
tools=["read_file", "grep", "glob", "git_diff", "git_log"],
|
|
44
|
+
system_prompt_template="""You are a planning subagent for SKY.
|
|
45
|
+
|
|
46
|
+
Your role is to analyze the codebase and create a detailed plan.
|
|
47
|
+
|
|
48
|
+
Available tools: read_file, grep, glob, git_diff, git_log
|
|
49
|
+
|
|
50
|
+
Output a plan with:
|
|
51
|
+
1. What needs to be done
|
|
52
|
+
2. What files will be affected
|
|
53
|
+
3. Potential risks or dependencies
|
|
54
|
+
4. Estimated effort
|
|
55
|
+
|
|
56
|
+
Be thorough and specific.""",
|
|
57
|
+
max_turns=5,
|
|
58
|
+
temperature=0.3,
|
|
59
|
+
),
|
|
60
|
+
SubagentRole.CODER: SubagentConfig(
|
|
61
|
+
role=SubagentRole.CODER,
|
|
62
|
+
model_role="fast_loop",
|
|
63
|
+
tools=["read_file", "write_file", "edit_file", "grep", "glob", "git_diff", "git_commit", "git_branch"],
|
|
64
|
+
system_prompt_template="""You are a coding subagent for SKY.
|
|
65
|
+
|
|
66
|
+
Your role is to implement code changes.
|
|
67
|
+
|
|
68
|
+
Available tools: read_file, write_file, edit_file, grep, glob, git_diff, git_commit, git_branch
|
|
69
|
+
|
|
70
|
+
Guidelines:
|
|
71
|
+
1. Read relevant files first
|
|
72
|
+
2. Make targeted changes
|
|
73
|
+
3. Use git_commit to commit changes
|
|
74
|
+
4. Explain what you changed and why
|
|
75
|
+
|
|
76
|
+
Be precise and focused.""",
|
|
77
|
+
max_turns=10,
|
|
78
|
+
temperature=0.1,
|
|
79
|
+
),
|
|
80
|
+
SubagentRole.TESTER: SubagentConfig(
|
|
81
|
+
role=SubagentRole.TESTER,
|
|
82
|
+
model_role="fast_loop",
|
|
83
|
+
tools=["read_file", "run_tests", "lint", "grep", "glob"],
|
|
84
|
+
system_prompt_template="""You are a testing subagent for SKY.
|
|
85
|
+
|
|
86
|
+
Your role is to run tests and report results.
|
|
87
|
+
|
|
88
|
+
Available tools: read_file, run_tests, lint, grep, glob
|
|
89
|
+
|
|
90
|
+
Guidelines:
|
|
91
|
+
1. Run tests after code changes
|
|
92
|
+
2. Report pass/fail results
|
|
93
|
+
3. If tests fail, provide error details
|
|
94
|
+
4. Suggest fixes for failing tests
|
|
95
|
+
|
|
96
|
+
Be thorough and accurate.""",
|
|
97
|
+
max_turns=5,
|
|
98
|
+
temperature=0.1,
|
|
99
|
+
),
|
|
100
|
+
SubagentRole.REVIEWER: SubagentConfig(
|
|
101
|
+
role=SubagentRole.REVIEWER,
|
|
102
|
+
model_role="planning",
|
|
103
|
+
tools=["read_file", "grep", "glob", "git_diff", "git_log"],
|
|
104
|
+
system_prompt_template="""You are a code reviewer subagent for SKY.
|
|
105
|
+
|
|
106
|
+
Your role is to review changes and provide feedback.
|
|
107
|
+
|
|
108
|
+
Available tools: read_file, grep, glob, git_diff, git_log
|
|
109
|
+
|
|
110
|
+
Guidelines:
|
|
111
|
+
1. Review the diff carefully
|
|
112
|
+
2. Check for code quality issues
|
|
113
|
+
3. Look for potential bugs or edge cases
|
|
114
|
+
4. Provide constructive feedback
|
|
115
|
+
5. Suggest improvements
|
|
116
|
+
|
|
117
|
+
Be critical but constructive.""",
|
|
118
|
+
max_turns=5,
|
|
119
|
+
temperature=0.3,
|
|
120
|
+
),
|
|
121
|
+
}
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
class SubagentRunner:
|
|
125
|
+
"""Runner for subagent execution."""
|
|
126
|
+
|
|
127
|
+
def __init__(
|
|
128
|
+
self,
|
|
129
|
+
config: DexProjectConfig,
|
|
130
|
+
db: DatabaseManager,
|
|
131
|
+
router: ModelRouter,
|
|
132
|
+
approval_gate: ApprovalGate,
|
|
133
|
+
indexer: Optional[RepoIndexer] = None,
|
|
134
|
+
):
|
|
135
|
+
self.config = config
|
|
136
|
+
self.db = db
|
|
137
|
+
self.router = router
|
|
138
|
+
self.approval_gate = approval_gate
|
|
139
|
+
self.indexer = indexer
|
|
140
|
+
|
|
141
|
+
async def run(
|
|
142
|
+
self,
|
|
143
|
+
task: Dict[str, Any],
|
|
144
|
+
role: SubagentRole,
|
|
145
|
+
parent_state: Optional[Dict[str, Any]] = None,
|
|
146
|
+
event_queue: Optional[Any] = None,
|
|
147
|
+
) -> Dict[str, Any]:
|
|
148
|
+
"""Run a subagent for a specific task."""
|
|
149
|
+
logger.info(f"Running subagent {role} for task {task.get('id', 'unknown')}")
|
|
150
|
+
|
|
151
|
+
if event_queue:
|
|
152
|
+
import asyncio
|
|
153
|
+
if isinstance(event_queue, asyncio.Queue):
|
|
154
|
+
await event_queue.put({
|
|
155
|
+
"type": "subagent_start",
|
|
156
|
+
"role": role.value,
|
|
157
|
+
"task": task.get("description", "Unknown task")
|
|
158
|
+
})
|
|
159
|
+
|
|
160
|
+
# Get role configuration
|
|
161
|
+
role_config = ROLE_CONFIGS.get(role)
|
|
162
|
+
if not role_config:
|
|
163
|
+
raise ValueError(f"Unknown role: {role}")
|
|
164
|
+
|
|
165
|
+
# Build system prompt
|
|
166
|
+
system_prompt = role_config.system_prompt_template
|
|
167
|
+
|
|
168
|
+
# Add task-specific context
|
|
169
|
+
if parent_state:
|
|
170
|
+
context = self._build_context(parent_state)
|
|
171
|
+
system_prompt = f"{system_prompt}\n\nContext from parent:\n{context}"
|
|
172
|
+
|
|
173
|
+
# Build messages
|
|
174
|
+
messages = [
|
|
175
|
+
{"role": "system", "content": system_prompt},
|
|
176
|
+
{"role": "user", "content": f"Task: {task.get('description', '')}\n\nDependencies: {task.get('dependencies', [])}\n\nRequired tools: {task.get('tools_needed', [])}"}
|
|
177
|
+
]
|
|
178
|
+
|
|
179
|
+
# Create a fresh session for subagent
|
|
180
|
+
session_id = self.db.create_session(f"subagent_{role.value}", task.get("id", "unknown"))
|
|
181
|
+
|
|
182
|
+
# Initialize FastLoopEngine for this subagent
|
|
183
|
+
engine = FastLoopEngine(
|
|
184
|
+
config=self.config,
|
|
185
|
+
db=self.db,
|
|
186
|
+
router=self.router,
|
|
187
|
+
approval_gate=self.approval_gate,
|
|
188
|
+
session_id=session_id,
|
|
189
|
+
indexer=self.indexer,
|
|
190
|
+
)
|
|
191
|
+
|
|
192
|
+
# Run the fast loop
|
|
193
|
+
tool_filter = role_config.tools if role_config.tools else None
|
|
194
|
+
|
|
195
|
+
responses = []
|
|
196
|
+
try:
|
|
197
|
+
async for event in engine.run(
|
|
198
|
+
messages=messages,
|
|
199
|
+
mode="agent",
|
|
200
|
+
max_turns=role_config.max_turns,
|
|
201
|
+
tools_filter=tool_filter,
|
|
202
|
+
inject_context=False, # Subagents don't auto-inject context
|
|
203
|
+
):
|
|
204
|
+
if event_queue:
|
|
205
|
+
import asyncio
|
|
206
|
+
if isinstance(event_queue, asyncio.Queue):
|
|
207
|
+
if event.get("type") in ("tool_call", "tool_result", "error"):
|
|
208
|
+
await event_queue.put(event)
|
|
209
|
+
|
|
210
|
+
if event.get("type") == "final_answer":
|
|
211
|
+
responses.append(event.get("content", ""))
|
|
212
|
+
elif event.get("type") == "error":
|
|
213
|
+
error_msg = event.get("content", event.get("error", "Unknown error"))
|
|
214
|
+
if event_queue:
|
|
215
|
+
import asyncio
|
|
216
|
+
if isinstance(event_queue, asyncio.Queue):
|
|
217
|
+
await event_queue.put({
|
|
218
|
+
"type": "subagent_complete",
|
|
219
|
+
"role": role.value,
|
|
220
|
+
"summary": f"Error: {error_msg}"
|
|
221
|
+
})
|
|
222
|
+
return {
|
|
223
|
+
"task_id": task.get("id"),
|
|
224
|
+
"success": False,
|
|
225
|
+
"error": error_msg,
|
|
226
|
+
"summary": "",
|
|
227
|
+
}
|
|
228
|
+
except Exception as e:
|
|
229
|
+
logger.error(f"Subagent {role} failed: {e}")
|
|
230
|
+
if event_queue:
|
|
231
|
+
import asyncio
|
|
232
|
+
if isinstance(event_queue, asyncio.Queue):
|
|
233
|
+
await event_queue.put({
|
|
234
|
+
"type": "subagent_complete",
|
|
235
|
+
"role": role.value,
|
|
236
|
+
"summary": f"Failed: {str(e)}"
|
|
237
|
+
})
|
|
238
|
+
return {
|
|
239
|
+
"task_id": task.get("id"),
|
|
240
|
+
"success": False,
|
|
241
|
+
"error": str(e),
|
|
242
|
+
"summary": "",
|
|
243
|
+
}
|
|
244
|
+
|
|
245
|
+
# Compress results
|
|
246
|
+
summary = self._summarize_results(responses, role)
|
|
247
|
+
|
|
248
|
+
if event_queue:
|
|
249
|
+
import asyncio
|
|
250
|
+
if isinstance(event_queue, asyncio.Queue):
|
|
251
|
+
await event_queue.put({
|
|
252
|
+
"type": "subagent_complete",
|
|
253
|
+
"role": role.value,
|
|
254
|
+
"summary": summary
|
|
255
|
+
})
|
|
256
|
+
|
|
257
|
+
return {
|
|
258
|
+
"task_id": task.get("id"),
|
|
259
|
+
"role": role.value,
|
|
260
|
+
"success": True,
|
|
261
|
+
"responses": responses,
|
|
262
|
+
"summary": summary,
|
|
263
|
+
"session_id": session_id,
|
|
264
|
+
"tool_usage": engine._tool_usage if hasattr(engine, "_tool_usage") else [],
|
|
265
|
+
}
|
|
266
|
+
|
|
267
|
+
def _build_context(self, parent_state: Dict[str, Any]) -> str:
|
|
268
|
+
"""Build context from parent state."""
|
|
269
|
+
context_parts = []
|
|
270
|
+
|
|
271
|
+
if parent_state.get("refined_goal"):
|
|
272
|
+
context_parts.append(f"Goal: {parent_state['refined_goal']}")
|
|
273
|
+
|
|
274
|
+
if parent_state.get("task_results"):
|
|
275
|
+
completed = [
|
|
276
|
+
f"{tid}: {res.get('summary', 'completed')}"
|
|
277
|
+
for tid, res in parent_state["task_results"].items()
|
|
278
|
+
if res.get("success")
|
|
279
|
+
]
|
|
280
|
+
if completed:
|
|
281
|
+
context_parts.append(f"Completed tasks:\n- " + "\n- ".join(completed))
|
|
282
|
+
|
|
283
|
+
return "\n".join(context_parts) if context_parts else "No additional context"
|
|
284
|
+
|
|
285
|
+
def _summarize_results(self, responses: List[str], role: SubagentRole) -> str:
|
|
286
|
+
"""Summarize subagent results."""
|
|
287
|
+
if not responses:
|
|
288
|
+
return "No output from subagent"
|
|
289
|
+
|
|
290
|
+
# For testers, summarize test results
|
|
291
|
+
if role == SubagentRole.TESTER:
|
|
292
|
+
# Look for test results
|
|
293
|
+
for response in responses:
|
|
294
|
+
if "passed" in response.lower() or "failed" in response.lower():
|
|
295
|
+
# Extract test summary
|
|
296
|
+
lines = response.split("\n")
|
|
297
|
+
test_lines = [l for l in lines if "pass" in l.lower() or "fail" in l.lower() or "test" in l.lower()]
|
|
298
|
+
if test_lines:
|
|
299
|
+
return "\n".join(test_lines[:5])
|
|
300
|
+
|
|
301
|
+
# Default: return first response truncated
|
|
302
|
+
return responses[-1][:500] + ("..." if len(responses[-1]) > 500 else "")
|