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/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 "")