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/workflow.py
ADDED
|
@@ -0,0 +1,662 @@
|
|
|
1
|
+
"""LangGraph multi-step workflow engine for SKY."""
|
|
2
|
+
|
|
3
|
+
from typing import TypedDict, List, Dict, Any, Optional, Literal, AsyncIterator
|
|
4
|
+
from dataclasses import dataclass, field
|
|
5
|
+
from enum import Enum
|
|
6
|
+
import json
|
|
7
|
+
import logging
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
import asyncio
|
|
10
|
+
|
|
11
|
+
from langgraph.graph import StateGraph, END
|
|
12
|
+
from langgraph.checkpoint.sqlite import SqliteSaver
|
|
13
|
+
from langgraph.constants import START
|
|
14
|
+
|
|
15
|
+
from sky.config.schema import DexProjectConfig
|
|
16
|
+
from sky.storage.db import DatabaseManager
|
|
17
|
+
from sky.core.router import ModelRouter
|
|
18
|
+
from sky.core.approval import ApprovalGate
|
|
19
|
+
from sky.core.fast_loop import FastLoopEngine
|
|
20
|
+
from sky.core.subagent import SubagentRunner, SubagentRole
|
|
21
|
+
from sky.memory.indexer import RepoIndexer
|
|
22
|
+
|
|
23
|
+
logger = logging.getLogger(__name__)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class WorkflowStatus(str, Enum):
|
|
27
|
+
"""Workflow execution status."""
|
|
28
|
+
PENDING = "pending"
|
|
29
|
+
RUNNING = "running"
|
|
30
|
+
REFINING = "refining"
|
|
31
|
+
DECOMPOSING = "decomposing"
|
|
32
|
+
EXECUTING = "executing"
|
|
33
|
+
TESTING = "testing"
|
|
34
|
+
REFLECTING = "reflecting"
|
|
35
|
+
COMPLETED = "completed"
|
|
36
|
+
FAILED = "failed"
|
|
37
|
+
CANCELLED = "cancelled"
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class WorkflowState(TypedDict):
|
|
41
|
+
"""State for the LangGraph workflow."""
|
|
42
|
+
# Input
|
|
43
|
+
user_goal: str
|
|
44
|
+
mode: str
|
|
45
|
+
|
|
46
|
+
# Refinement
|
|
47
|
+
refined_goal: str
|
|
48
|
+
goal_confidence: float
|
|
49
|
+
|
|
50
|
+
# Decomposition
|
|
51
|
+
tasks: List[Dict[str, Any]] # [{"id", "description", "dependencies", "estimated_effort", "tools_needed", "role"}]
|
|
52
|
+
task_dag: Dict[str, List[str]] # task_id -> [dependent_task_ids]
|
|
53
|
+
|
|
54
|
+
# Execution
|
|
55
|
+
current_task_index: int
|
|
56
|
+
task_results: Dict[str, Any] # task_id -> result
|
|
57
|
+
task_status: Dict[str, str] # task_id -> "pending|running|completed|failed"
|
|
58
|
+
|
|
59
|
+
# Subagent
|
|
60
|
+
subagent_results: List[Dict[str, Any]]
|
|
61
|
+
|
|
62
|
+
# Testing
|
|
63
|
+
test_results: Dict[str, Any] # {"passed": int, "failed": int, "errors": List[str]}
|
|
64
|
+
retry_count: int
|
|
65
|
+
max_retries: int
|
|
66
|
+
|
|
67
|
+
# Output
|
|
68
|
+
diff: str
|
|
69
|
+
final_summary: str
|
|
70
|
+
status: WorkflowStatus
|
|
71
|
+
|
|
72
|
+
# Metadata
|
|
73
|
+
session_id: str
|
|
74
|
+
turn_count: int
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
@dataclass
|
|
78
|
+
class WorkflowConfig:
|
|
79
|
+
"""Configuration for workflow execution."""
|
|
80
|
+
max_retries: int = 3
|
|
81
|
+
max_subtasks: int = 20
|
|
82
|
+
parallel_limit: int = 3
|
|
83
|
+
enable_subagents: bool = True
|
|
84
|
+
enable_test_and_reflect: bool = True
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class WorkflowEngine:
|
|
88
|
+
"""LangGraph-based workflow engine for multi-step SDLC."""
|
|
89
|
+
|
|
90
|
+
def __init__(
|
|
91
|
+
self,
|
|
92
|
+
config: DexProjectConfig,
|
|
93
|
+
db: DatabaseManager,
|
|
94
|
+
router: ModelRouter,
|
|
95
|
+
approval_gate: ApprovalGate,
|
|
96
|
+
indexer: Optional[RepoIndexer] = None,
|
|
97
|
+
checkpoint_dir: Optional[Path] = None,
|
|
98
|
+
):
|
|
99
|
+
self.config = config
|
|
100
|
+
self.db = db
|
|
101
|
+
self.router = router
|
|
102
|
+
self.approval_gate = approval_gate
|
|
103
|
+
self.indexer = indexer
|
|
104
|
+
self.checkpoint_dir = checkpoint_dir or Path.cwd() / ".sky" / "checkpoints"
|
|
105
|
+
self.checkpoint_dir.mkdir(parents=True, exist_ok=True)
|
|
106
|
+
|
|
107
|
+
self.workflow_config = WorkflowConfig(
|
|
108
|
+
max_retries=getattr(config, "workflow_max_retries", 3),
|
|
109
|
+
parallel_limit=getattr(config, "workflow_parallel_limit", 3)
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
self._graph = None
|
|
113
|
+
self._checkpointer = None
|
|
114
|
+
self._subagent_runner = None
|
|
115
|
+
|
|
116
|
+
def _build_graph(self) -> StateGraph:
|
|
117
|
+
"""Build the LangGraph workflow graph."""
|
|
118
|
+
graph = StateGraph(WorkflowState)
|
|
119
|
+
|
|
120
|
+
# Add nodes
|
|
121
|
+
graph.add_node("refine_goal", self._refine_goal)
|
|
122
|
+
graph.add_node("decompose", self._decompose)
|
|
123
|
+
graph.add_node("build_dag", self._build_dag)
|
|
124
|
+
graph.add_node("dispatch_subagents", self._dispatch_subagents)
|
|
125
|
+
graph.add_node("run_tests", self._run_tests)
|
|
126
|
+
graph.add_node("reflect", self._reflect)
|
|
127
|
+
graph.add_node("summarize", self._summarize)
|
|
128
|
+
|
|
129
|
+
# Add edges
|
|
130
|
+
graph.add_edge(START, "refine_goal")
|
|
131
|
+
graph.add_edge("refine_goal", "decompose")
|
|
132
|
+
graph.add_edge("decompose", "build_dag")
|
|
133
|
+
graph.add_edge("build_dag", "dispatch_subagents")
|
|
134
|
+
graph.add_edge("dispatch_subagents", "run_tests")
|
|
135
|
+
|
|
136
|
+
# Conditional edge for reflect
|
|
137
|
+
graph.add_conditional_edges(
|
|
138
|
+
"reflect",
|
|
139
|
+
self._should_retry,
|
|
140
|
+
{
|
|
141
|
+
"retry": "dispatch_subagents",
|
|
142
|
+
"summarize": "summarize",
|
|
143
|
+
"fail": "summarize",
|
|
144
|
+
}
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
graph.add_edge("run_tests", "reflect")
|
|
148
|
+
graph.add_edge("summarize", END)
|
|
149
|
+
|
|
150
|
+
return graph
|
|
151
|
+
|
|
152
|
+
async def _emit(self, event: Dict[str, Any]) -> None:
|
|
153
|
+
if hasattr(self, "event_queue") and self.event_queue:
|
|
154
|
+
import asyncio
|
|
155
|
+
if isinstance(self.event_queue, asyncio.Queue):
|
|
156
|
+
await self.event_queue.put(event)
|
|
157
|
+
|
|
158
|
+
async def _refine_goal(self, state: WorkflowState) -> WorkflowState:
|
|
159
|
+
"""Refine the user's goal using the planning model."""
|
|
160
|
+
await self._emit({"type": "step", "step": "refining", "message": "🎯 Refining goal..."})
|
|
161
|
+
logger.info("Refining goal...")
|
|
162
|
+
|
|
163
|
+
messages = [
|
|
164
|
+
{"role": "system", "content": """You are a goal refinement expert. Take the user's raw request and refine it into a clear, actionable software development goal.
|
|
165
|
+
|
|
166
|
+
Output should be a JSON object with:
|
|
167
|
+
- refined_goal: A clear, specific description of what needs to be done
|
|
168
|
+
- confidence: A number between 0 and 1 indicating how confident you are in understanding the goal
|
|
169
|
+
- questions: Any clarifying questions (if confidence < 0.7)
|
|
170
|
+
|
|
171
|
+
Be specific about what files or components might be involved."""},
|
|
172
|
+
{"role": "user", "content": state["user_goal"]}
|
|
173
|
+
]
|
|
174
|
+
|
|
175
|
+
response, was_fallback, model_used = await self.router.route(
|
|
176
|
+
"planning",
|
|
177
|
+
messages,
|
|
178
|
+
expected_tool_calls=0,
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
content = response.get("content", "")
|
|
182
|
+
|
|
183
|
+
# Parse JSON response
|
|
184
|
+
try:
|
|
185
|
+
# Extract JSON from markdown if needed
|
|
186
|
+
import re
|
|
187
|
+
json_match = re.search(r'\{.*\}', content, re.DOTALL)
|
|
188
|
+
if json_match:
|
|
189
|
+
data = json.loads(json_match.group())
|
|
190
|
+
state["refined_goal"] = data.get("refined_goal", state["user_goal"])
|
|
191
|
+
state["goal_confidence"] = data.get("confidence", 0.5)
|
|
192
|
+
else:
|
|
193
|
+
state["refined_goal"] = content
|
|
194
|
+
state["goal_confidence"] = 0.5
|
|
195
|
+
except json.JSONDecodeError:
|
|
196
|
+
state["refined_goal"] = content
|
|
197
|
+
state["goal_confidence"] = 0.5
|
|
198
|
+
|
|
199
|
+
state["status"] = WorkflowStatus.REFINING.value
|
|
200
|
+
await self._emit({"type": "step_complete", "step": "refining", "result": state["refined_goal"]})
|
|
201
|
+
return state
|
|
202
|
+
|
|
203
|
+
async def _decompose(self, state: WorkflowState) -> WorkflowState:
|
|
204
|
+
"""Decompose the goal into subtasks."""
|
|
205
|
+
await self._emit({"type": "step", "step": "decomposing", "message": "Breaking down into tasks..."})
|
|
206
|
+
logger.info("Decomposing goal into tasks...")
|
|
207
|
+
|
|
208
|
+
messages = [
|
|
209
|
+
{"role": "system", "content": """You are a task decomposition expert. Break down the goal into a structured set of subtasks.
|
|
210
|
+
|
|
211
|
+
Output should be a JSON array of tasks, each with:
|
|
212
|
+
- id: Unique identifier (e.g., "task_1", "task_2")
|
|
213
|
+
- description: Clear description of what to do
|
|
214
|
+
- dependencies: List of task IDs that must complete first
|
|
215
|
+
- estimated_effort: "small", "medium", or "large"
|
|
216
|
+
- tools_needed: List of tools required (e.g., ["read_file", "edit_file"])
|
|
217
|
+
- role: "planner", "coder", "tester", "reviewer"
|
|
218
|
+
|
|
219
|
+
Example:
|
|
220
|
+
[
|
|
221
|
+
{"id": "task_1", "description": "Implement the requested feature", "dependencies": [], "estimated_effort": "medium", "tools_needed": ["read_file", "write_file", "edit_file"], "role": "coder"}
|
|
222
|
+
]
|
|
223
|
+
|
|
224
|
+
CRITICAL RULES:
|
|
225
|
+
1. For simple goals (like writing a single script or making a basic modification), use exactly ONE 'coder' task.
|
|
226
|
+
2. DO NOT over-complicate simple requests by splitting them into 'planning', 'setup', and 'coding' phases.
|
|
227
|
+
3. Only use multiple tasks for complex features that require separate architectural planning or extensive multi-file changes.
|
|
228
|
+
4. Keep tasks atomic and focused. Max 10 tasks."""},
|
|
229
|
+
{"role": "user", "content": f"Goal: {state.get('refined_goal', state['user_goal'])}"}
|
|
230
|
+
]
|
|
231
|
+
|
|
232
|
+
response, was_fallback, model_used = await self.router.route(
|
|
233
|
+
"planning",
|
|
234
|
+
messages,
|
|
235
|
+
expected_tool_calls=0,
|
|
236
|
+
)
|
|
237
|
+
|
|
238
|
+
content = response.get("content", "")
|
|
239
|
+
|
|
240
|
+
# Parse JSON response
|
|
241
|
+
try:
|
|
242
|
+
import re
|
|
243
|
+
json_match = re.search(r'\[.*\]', content, re.DOTALL)
|
|
244
|
+
if json_match:
|
|
245
|
+
tasks = json.loads(json_match.group())
|
|
246
|
+
state["tasks"] = tasks
|
|
247
|
+
else:
|
|
248
|
+
# Fallback: create single task
|
|
249
|
+
state["tasks"] = [{
|
|
250
|
+
"id": "task_1",
|
|
251
|
+
"description": state.get("refined_goal", state["user_goal"]),
|
|
252
|
+
"dependencies": [],
|
|
253
|
+
"estimated_effort": "medium",
|
|
254
|
+
"tools_needed": ["read_file", "write_file", "edit_file", "run_tests"],
|
|
255
|
+
"role": "coder"
|
|
256
|
+
}]
|
|
257
|
+
except json.JSONDecodeError:
|
|
258
|
+
state["tasks"] = [{
|
|
259
|
+
"id": "task_1",
|
|
260
|
+
"description": state.get("refined_goal", state["user_goal"]),
|
|
261
|
+
"dependencies": [],
|
|
262
|
+
"estimated_effort": "medium",
|
|
263
|
+
"tools_needed": ["read_file", "write_file", "edit_file", "run_tests"],
|
|
264
|
+
"role": "coder"
|
|
265
|
+
}]
|
|
266
|
+
|
|
267
|
+
# Initialize task status
|
|
268
|
+
state["task_status"] = {t["id"]: "pending" for t in state["tasks"]}
|
|
269
|
+
state["task_results"] = {}
|
|
270
|
+
if "subagent_results" not in state or not state["subagent_results"]:
|
|
271
|
+
state["subagent_results"] = []
|
|
272
|
+
state["status"] = WorkflowStatus.DECOMPOSING.value
|
|
273
|
+
|
|
274
|
+
return state
|
|
275
|
+
|
|
276
|
+
async def _build_dag(self, state: WorkflowState) -> WorkflowState:
|
|
277
|
+
"""Build the DAG from task dependencies."""
|
|
278
|
+
logger.info("Building DAG...")
|
|
279
|
+
|
|
280
|
+
dag = {}
|
|
281
|
+
for task in state["tasks"]:
|
|
282
|
+
dag[task["id"]] = task.get("dependencies", [])
|
|
283
|
+
|
|
284
|
+
state["task_dag"] = dag
|
|
285
|
+
state["status"] = WorkflowStatus.EXECUTING.value
|
|
286
|
+
|
|
287
|
+
return state
|
|
288
|
+
|
|
289
|
+
async def _dispatch_subagents(self, state: WorkflowState) -> WorkflowState:
|
|
290
|
+
"""Dispatch subagents for parallel task execution."""
|
|
291
|
+
await self._emit({"type": "step", "step": "executing", "message": "🔧 Executing subagents..."})
|
|
292
|
+
logger.info("Dispatching subagents...")
|
|
293
|
+
|
|
294
|
+
# Get pending tasks
|
|
295
|
+
pending_tasks = [
|
|
296
|
+
t for t in state["tasks"]
|
|
297
|
+
if state["task_status"].get(t["id"]) == "pending"
|
|
298
|
+
]
|
|
299
|
+
|
|
300
|
+
if not pending_tasks:
|
|
301
|
+
return state
|
|
302
|
+
|
|
303
|
+
# Check dependencies
|
|
304
|
+
available_tasks = []
|
|
305
|
+
for task in pending_tasks:
|
|
306
|
+
deps = task.get("dependencies", [])
|
|
307
|
+
if all(state["task_status"].get(d) == "completed" for d in deps):
|
|
308
|
+
available_tasks.append(task)
|
|
309
|
+
|
|
310
|
+
if not available_tasks:
|
|
311
|
+
# No tasks available - wait for dependencies
|
|
312
|
+
return state
|
|
313
|
+
|
|
314
|
+
# Limit parallel execution
|
|
315
|
+
available_tasks = available_tasks[:self.workflow_config.parallel_limit]
|
|
316
|
+
|
|
317
|
+
# Create subagent runner if not exists
|
|
318
|
+
if self._subagent_runner is None:
|
|
319
|
+
self._subagent_runner = SubagentRunner(
|
|
320
|
+
config=self.config,
|
|
321
|
+
db=self.db,
|
|
322
|
+
router=self.router,
|
|
323
|
+
approval_gate=self.approval_gate,
|
|
324
|
+
indexer=self.indexer,
|
|
325
|
+
)
|
|
326
|
+
|
|
327
|
+
# Execute tasks in parallel using asyncio.gather
|
|
328
|
+
tasks_coros = []
|
|
329
|
+
for task in available_tasks:
|
|
330
|
+
# Mark as running
|
|
331
|
+
state["task_status"][task["id"]] = "running"
|
|
332
|
+
|
|
333
|
+
# Get role
|
|
334
|
+
role = SubagentRole(task.get("role", "coder"))
|
|
335
|
+
|
|
336
|
+
# Emit task_start before running
|
|
337
|
+
await self._emit({"type": "task_start", "task": task["id"], "message": f"🔧 Executing {task['id']}..."})
|
|
338
|
+
|
|
339
|
+
# Create coroutine
|
|
340
|
+
coro = self._subagent_runner.run(
|
|
341
|
+
task=task,
|
|
342
|
+
role=role,
|
|
343
|
+
parent_state=state,
|
|
344
|
+
event_queue=getattr(self, "event_queue", None)
|
|
345
|
+
)
|
|
346
|
+
tasks_coros.append(coro)
|
|
347
|
+
|
|
348
|
+
results = await asyncio.gather(*tasks_coros)
|
|
349
|
+
|
|
350
|
+
# Update state
|
|
351
|
+
for result in results:
|
|
352
|
+
task_id = result.get("task_id")
|
|
353
|
+
await self._emit({"type": "task_result", "task": task_id, "result": result})
|
|
354
|
+
if not result.get("success"):
|
|
355
|
+
await self._emit({"type": "error", "content": f"Task {task_id} failed: {result.get('error')}"})
|
|
356
|
+
state["task_status"][task_id] = "completed" if result.get("success") else "failed"
|
|
357
|
+
state["task_results"][task_id] = result
|
|
358
|
+
state["subagent_results"].append(result)
|
|
359
|
+
|
|
360
|
+
state["status"] = WorkflowStatus.EXECUTING.value
|
|
361
|
+
|
|
362
|
+
return state
|
|
363
|
+
|
|
364
|
+
async def _run_tests(self, state: WorkflowState) -> WorkflowState:
|
|
365
|
+
"""Run tests after code changes."""
|
|
366
|
+
await self._emit({"type": "step", "step": "testing", "message": "Running tests..."})
|
|
367
|
+
logger.info("Running tests...")
|
|
368
|
+
|
|
369
|
+
# Collect any test-related changes
|
|
370
|
+
test_files = []
|
|
371
|
+
for result in state.get("subagent_results", []):
|
|
372
|
+
if result.get("test_results"):
|
|
373
|
+
test_files.extend(result.get("test_files", []))
|
|
374
|
+
|
|
375
|
+
# Run tests using the test tool
|
|
376
|
+
from sky.tools.test_tools import run_tests
|
|
377
|
+
|
|
378
|
+
try:
|
|
379
|
+
# run_tests returns a dict matching exactly what we need
|
|
380
|
+
result = run_tests(timeout=120)
|
|
381
|
+
|
|
382
|
+
state["test_results"] = {
|
|
383
|
+
"passed": result.get("passed", 0),
|
|
384
|
+
"failed": result.get("failed", 0),
|
|
385
|
+
"errors": result.get("errors", []),
|
|
386
|
+
"success": result.get("success", False)
|
|
387
|
+
}
|
|
388
|
+
state["status"] = WorkflowStatus.TESTING.value
|
|
389
|
+
except Exception as e:
|
|
390
|
+
state["test_results"] = {
|
|
391
|
+
"passed": 0,
|
|
392
|
+
"failed": 1,
|
|
393
|
+
"errors": [str(e)],
|
|
394
|
+
"success": False
|
|
395
|
+
}
|
|
396
|
+
|
|
397
|
+
return state
|
|
398
|
+
|
|
399
|
+
async def _reflect(self, state: WorkflowState) -> WorkflowState:
|
|
400
|
+
"""Reflect on test results and decide next action."""
|
|
401
|
+
await self._emit({"type": "step", "step": "reflecting", "message": "Reflecting on results..."})
|
|
402
|
+
logger.info("Reflecting on results...")
|
|
403
|
+
|
|
404
|
+
test_results = state.get("test_results", {})
|
|
405
|
+
passed = test_results.get("passed", 0)
|
|
406
|
+
failed = test_results.get("failed", 0)
|
|
407
|
+
errors = test_results.get("errors", [])
|
|
408
|
+
|
|
409
|
+
if passed > 0 and failed == 0 and not errors:
|
|
410
|
+
# All tests passed
|
|
411
|
+
state["status"] = WorkflowStatus.COMPLETED.value
|
|
412
|
+
return state
|
|
413
|
+
|
|
414
|
+
if passed == 0 and failed == 0 and not errors:
|
|
415
|
+
# No tests ran or found, but check if all tasks succeeded
|
|
416
|
+
if all(status == "completed" for status in state["task_status"].values()):
|
|
417
|
+
state["status"] = WorkflowStatus.COMPLETED.value
|
|
418
|
+
return state
|
|
419
|
+
|
|
420
|
+
# Check retry count
|
|
421
|
+
retry_count = state.get("retry_count", 0)
|
|
422
|
+
max_retries = state.get("max_retries", self.workflow_config.max_retries)
|
|
423
|
+
|
|
424
|
+
if retry_count >= max_retries:
|
|
425
|
+
# Max retries exceeded
|
|
426
|
+
state["status"] = WorkflowStatus.FAILED.value
|
|
427
|
+
return state
|
|
428
|
+
|
|
429
|
+
await self._emit({"type": "step", "step": "retry", "message": f"Retrying tasks (attempt {retry_count + 1} of {max_retries})..."})
|
|
430
|
+
|
|
431
|
+
# Prepare for retry
|
|
432
|
+
state["retry_count"] = retry_count + 1
|
|
433
|
+
state["status"] = WorkflowStatus.REFLECTING.value
|
|
434
|
+
|
|
435
|
+
# Generate reflection message
|
|
436
|
+
messages = [
|
|
437
|
+
{"role": "system", "content": """You are a code reviewer. Analyze the test failures and suggest fixes.
|
|
438
|
+
|
|
439
|
+
Output a JSON with:
|
|
440
|
+
- analysis: Brief analysis of what went wrong
|
|
441
|
+
- suggested_fix: Description of what to change
|
|
442
|
+
- files_to_modify: List of file paths that need changes"""},
|
|
443
|
+
{"role": "user", "content": f"""
|
|
444
|
+
Test Results:
|
|
445
|
+
- Passed: {passed}
|
|
446
|
+
- Failed: {failed}
|
|
447
|
+
- Errors: {errors}
|
|
448
|
+
|
|
449
|
+
Tasks completed: {[t for t in state['task_status'] if state['task_status'][t] == 'completed']}
|
|
450
|
+
Failed tasks: {[t for t in state['task_status'] if state['task_status'][t] == 'failed']}
|
|
451
|
+
|
|
452
|
+
Provide analysis and suggested fixes."""}
|
|
453
|
+
]
|
|
454
|
+
|
|
455
|
+
response, was_fallback, model_used = await self.router.route(
|
|
456
|
+
"planning",
|
|
457
|
+
messages,
|
|
458
|
+
expected_tool_calls=0,
|
|
459
|
+
)
|
|
460
|
+
|
|
461
|
+
content = response.get("content", "")
|
|
462
|
+
|
|
463
|
+
# Parse reflection
|
|
464
|
+
try:
|
|
465
|
+
import re
|
|
466
|
+
json_match = re.search(r'\{.*\}', content, re.DOTALL)
|
|
467
|
+
if json_match:
|
|
468
|
+
reflection = json.loads(json_match.group())
|
|
469
|
+
# Add failed tasks back to pending
|
|
470
|
+
for task_id, status in state["task_status"].items():
|
|
471
|
+
if status == "failed":
|
|
472
|
+
state["task_status"][task_id] = "pending"
|
|
473
|
+
else:
|
|
474
|
+
# Fallback: retry failed tasks
|
|
475
|
+
for task_id, status in state["task_status"].items():
|
|
476
|
+
if status == "failed":
|
|
477
|
+
state["task_status"][task_id] = "pending"
|
|
478
|
+
except json.JSONDecodeError:
|
|
479
|
+
# Fallback: retry failed tasks
|
|
480
|
+
for task_id, status in state["task_status"].items():
|
|
481
|
+
if status == "failed":
|
|
482
|
+
state["task_status"][task_id] = "pending"
|
|
483
|
+
|
|
484
|
+
return state
|
|
485
|
+
|
|
486
|
+
def _should_retry(self, state: WorkflowState) -> Literal["retry", "summarize", "fail"]:
|
|
487
|
+
"""Determine if workflow should retry."""
|
|
488
|
+
status = state.get("status")
|
|
489
|
+
|
|
490
|
+
if status == WorkflowStatus.COMPLETED.value:
|
|
491
|
+
return "summarize"
|
|
492
|
+
|
|
493
|
+
if status == WorkflowStatus.FAILED.value:
|
|
494
|
+
return "fail"
|
|
495
|
+
|
|
496
|
+
if status == WorkflowStatus.REFLECTING.value:
|
|
497
|
+
retry_count = state.get("retry_count", 0)
|
|
498
|
+
max_retries = state.get("max_retries", self.workflow_config.max_retries)
|
|
499
|
+
|
|
500
|
+
if retry_count < max_retries:
|
|
501
|
+
return "retry"
|
|
502
|
+
else:
|
|
503
|
+
return "fail"
|
|
504
|
+
|
|
505
|
+
return "summarize"
|
|
506
|
+
|
|
507
|
+
async def _summarize(self, state: WorkflowState) -> WorkflowState:
|
|
508
|
+
"""Generate final summary."""
|
|
509
|
+
await self._emit({"type": "step", "step": "summarizing", "message": "Generating summary..."})
|
|
510
|
+
logger.info("Generating summary...")
|
|
511
|
+
|
|
512
|
+
# Collect results
|
|
513
|
+
task_summary = []
|
|
514
|
+
for task in state.get("tasks", []):
|
|
515
|
+
task_id = task["id"]
|
|
516
|
+
status = state["task_status"].get(task_id, "unknown")
|
|
517
|
+
result = state["task_results"].get(task_id, {})
|
|
518
|
+
|
|
519
|
+
task_summary.append({
|
|
520
|
+
"id": task_id,
|
|
521
|
+
"description": task["description"],
|
|
522
|
+
"status": status,
|
|
523
|
+
"result": result.get("summary", ""),
|
|
524
|
+
})
|
|
525
|
+
|
|
526
|
+
test_results = state.get("test_results", {})
|
|
527
|
+
|
|
528
|
+
# Generate summary
|
|
529
|
+
messages = [
|
|
530
|
+
{"role": "system", "content": """You are a project summary expert. Create a clear, concise summary of the work completed.
|
|
531
|
+
|
|
532
|
+
Include:
|
|
533
|
+
1. What was accomplished
|
|
534
|
+
2. What was changed
|
|
535
|
+
3. Test results
|
|
536
|
+
4. Any issues or warnings
|
|
537
|
+
5. Next steps (if any)"""},
|
|
538
|
+
{"role": "user", "content": f"""
|
|
539
|
+
Original Goal: {state['user_goal']}
|
|
540
|
+
Refined Goal: {state.get('refined_goal', state['user_goal'])}
|
|
541
|
+
|
|
542
|
+
Task Summary: {json.dumps(task_summary, indent=2)}
|
|
543
|
+
|
|
544
|
+
Test Results: {json.dumps(test_results, indent=2)}
|
|
545
|
+
|
|
546
|
+
Create a summary of the work."""}
|
|
547
|
+
]
|
|
548
|
+
|
|
549
|
+
response, was_fallback, model_used = await self.router.route(
|
|
550
|
+
"planning",
|
|
551
|
+
messages,
|
|
552
|
+
expected_tool_calls=0,
|
|
553
|
+
)
|
|
554
|
+
|
|
555
|
+
content = response.get("content", "")
|
|
556
|
+
|
|
557
|
+
if state.get("status") != WorkflowStatus.FAILED.value:
|
|
558
|
+
state["status"] = WorkflowStatus.COMPLETED.value
|
|
559
|
+
|
|
560
|
+
await self._emit({"type": "complete", "summary": content})
|
|
561
|
+
|
|
562
|
+
return state
|
|
563
|
+
|
|
564
|
+
async def run_streaming(
|
|
565
|
+
self,
|
|
566
|
+
goal: str,
|
|
567
|
+
mode: str = "agent",
|
|
568
|
+
max_retries: int = 3,
|
|
569
|
+
resume_from: Optional[str] = None,
|
|
570
|
+
) -> AsyncIterator[Dict[str, Any]]:
|
|
571
|
+
"""Run the workflow with real-time event streaming."""
|
|
572
|
+
import asyncio
|
|
573
|
+
self.event_queue = asyncio.Queue()
|
|
574
|
+
|
|
575
|
+
# Start graph execution in background
|
|
576
|
+
task = asyncio.create_task(self.run(goal, mode, max_retries, resume_from))
|
|
577
|
+
|
|
578
|
+
while not task.done():
|
|
579
|
+
try:
|
|
580
|
+
event = await asyncio.wait_for(self.event_queue.get(), timeout=0.1)
|
|
581
|
+
yield event
|
|
582
|
+
if event.get("type") in ("complete", "error"):
|
|
583
|
+
break
|
|
584
|
+
except asyncio.TimeoutError:
|
|
585
|
+
continue
|
|
586
|
+
|
|
587
|
+
# Make sure we didn't miss any events after task completion
|
|
588
|
+
while not self.event_queue.empty():
|
|
589
|
+
event = await self.event_queue.get()
|
|
590
|
+
yield event
|
|
591
|
+
|
|
592
|
+
await task
|
|
593
|
+
|
|
594
|
+
async def run(
|
|
595
|
+
self,
|
|
596
|
+
goal: str,
|
|
597
|
+
mode: str = "agent",
|
|
598
|
+
max_retries: int = 3,
|
|
599
|
+
resume_from: Optional[str] = None,
|
|
600
|
+
) -> Dict[str, Any]:
|
|
601
|
+
"""Run the workflow."""
|
|
602
|
+
# Initialize state
|
|
603
|
+
state: WorkflowState = {
|
|
604
|
+
"user_goal": goal,
|
|
605
|
+
"mode": mode,
|
|
606
|
+
"refined_goal": "",
|
|
607
|
+
"goal_confidence": 0.0,
|
|
608
|
+
"tasks": [],
|
|
609
|
+
"task_dag": {},
|
|
610
|
+
"current_task_index": 0,
|
|
611
|
+
"task_results": {},
|
|
612
|
+
"task_status": {},
|
|
613
|
+
"subagent_results": [],
|
|
614
|
+
"test_results": {},
|
|
615
|
+
"retry_count": 0,
|
|
616
|
+
"max_retries": max_retries,
|
|
617
|
+
"diff": "",
|
|
618
|
+
"final_summary": "",
|
|
619
|
+
"status": WorkflowStatus.PENDING.value,
|
|
620
|
+
"session_id": self.db.create_session("workflow", goal),
|
|
621
|
+
"turn_count": 0,
|
|
622
|
+
}
|
|
623
|
+
|
|
624
|
+
# Build graph
|
|
625
|
+
self._graph = self._build_graph()
|
|
626
|
+
|
|
627
|
+
# Setup checkpointer
|
|
628
|
+
checkpoint_path = self.checkpoint_dir / f"{state['session_id']}.sqlite"
|
|
629
|
+
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
|
630
|
+
|
|
631
|
+
async with AsyncSqliteSaver.from_conn_string(str(checkpoint_path)) as checkpointer:
|
|
632
|
+
# Compile graph
|
|
633
|
+
app = self._graph.compile(checkpointer=checkpointer)
|
|
634
|
+
|
|
635
|
+
# Run workflow
|
|
636
|
+
config = {"configurable": {"thread_id": state["session_id"]}}
|
|
637
|
+
|
|
638
|
+
# If resuming
|
|
639
|
+
if resume_from:
|
|
640
|
+
snapshot = await app.aget_state(config)
|
|
641
|
+
if snapshot:
|
|
642
|
+
state = snapshot.values
|
|
643
|
+
|
|
644
|
+
# Execute asynchronously
|
|
645
|
+
async for event in app.astream(state, config):
|
|
646
|
+
pass
|
|
647
|
+
|
|
648
|
+
# Get final state
|
|
649
|
+
final_state_snapshot = await app.aget_state(config)
|
|
650
|
+
|
|
651
|
+
return final_state_snapshot.values
|
|
652
|
+
|
|
653
|
+
async def resume(self, session_id: str) -> Dict[str, Any]:
|
|
654
|
+
"""Resume an interrupted workflow."""
|
|
655
|
+
checkpoint_path = self.checkpoint_dir / f"{session_id}.sqlite"
|
|
656
|
+
if not checkpoint_path.exists():
|
|
657
|
+
raise ValueError(f"Session {session_id} not found")
|
|
658
|
+
|
|
659
|
+
return await self.run(
|
|
660
|
+
goal="Resume workflow",
|
|
661
|
+
resume_from=session_id,
|
|
662
|
+
)
|
sky/errors.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""Custom error hierarchy for Sky."""
|
|
2
|
+
|
|
3
|
+
class SkyError(Exception):
|
|
4
|
+
"""Base exception for Sky."""
|
|
5
|
+
def __init__(self, message: str, code: str = None, suggestion: str = None):
|
|
6
|
+
self.message = message
|
|
7
|
+
self.code = code
|
|
8
|
+
self.suggestion = suggestion
|
|
9
|
+
super().__init__(message)
|
|
10
|
+
|
|
11
|
+
class ConfigurationError(SkyError):
|
|
12
|
+
"""Configuration errors (missing keys, invalid config)."""
|
|
13
|
+
pass
|
|
14
|
+
|
|
15
|
+
class ProviderError(SkyError):
|
|
16
|
+
"""Provider errors (API down, rate limits)."""
|
|
17
|
+
pass
|
|
18
|
+
|
|
19
|
+
class ToolError(SkyError):
|
|
20
|
+
"""Tool execution errors."""
|
|
21
|
+
pass
|
|
22
|
+
|
|
23
|
+
class WorkflowError(SkyError):
|
|
24
|
+
"""Workflow execution errors."""
|
|
25
|
+
pass
|
|
26
|
+
|
|
27
|
+
class ValidationError(SkyError):
|
|
28
|
+
"""Input validation errors."""
|
|
29
|
+
pass
|