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/cli.py ADDED
@@ -0,0 +1,833 @@
1
+ # SKY is awesome
2
+ """SKY Command Line Interface."""
3
+
4
+ import asyncio
5
+ import os
6
+ import sys
7
+ from pathlib import Path
8
+
9
+ # Force UTF-8 encoding for stdout on Windows to prevent UnicodeEncodeError
10
+ if sys.platform == "win32":
11
+ try:
12
+ import ctypes
13
+ # Set Windows console to UTF-8 (65001)
14
+ ctypes.windll.kernel32.SetConsoleOutputCP(65001)
15
+ except Exception:
16
+ pass
17
+
18
+ if sys.stdout.encoding != "utf-8" and hasattr(sys.stdout, "reconfigure"):
19
+ sys.stdout.reconfigure(encoding="utf-8")
20
+
21
+ import typer
22
+ from rich.console import Console
23
+
24
+ from sky import __version__
25
+
26
+ app = typer.Typer(help="SKY - Local CLI-based agentic software development assistant")
27
+ console = Console()
28
+
29
+
30
+ def version_callback(value: bool) -> None:
31
+ if value:
32
+ console.print(f"SKY CLI Version: {__version__}")
33
+ raise typer.Exit()
34
+
35
+
36
+ @app.callback()
37
+ def main(
38
+ version: bool = typer.Option(None, "--version", callback=version_callback, is_eager=True, help="Show version."),
39
+ verbose: bool = typer.Option(False, "--verbose", "-v", help="Enable verbose output.")
40
+ ) -> None:
41
+ """
42
+ ☁️ Sky — Build without boundaries.
43
+
44
+ Commands:
45
+ chat Start conversational chat with Sky
46
+ ask Ask a question (read-only)
47
+ plan Create a structured plan
48
+ agent Run agent mode with tools
49
+ workflow Run multi-step workflow
50
+ init Interactive setup
51
+ index Index repository
52
+ context Semantic search
53
+ stats Show usage statistics
54
+ sessions List sessions
55
+ audit Show audit log
56
+ """
57
+ pass
58
+
59
+
60
+ async def _run_loop(mode: str, prompt: str, inject_context: bool = False, quiet: bool = False, verbose: bool = False, max_turns: int = 20):
61
+ from dotenv import load_dotenv
62
+ load_dotenv()
63
+
64
+ from sky.config import load_config, load_models_config
65
+ from sky.core.approval import ApprovalGate
66
+ from sky.core.fast_loop import FastLoopEngine
67
+ from sky.core.router import ModelRouter
68
+ from sky.storage import get_db
69
+ from sky.errors import SkyError
70
+
71
+ config = load_config()
72
+ if verbose:
73
+ config.verbose = True
74
+
75
+ if getattr(config, "security_enabled", True):
76
+ from sky.security.rate_limit import get_rate_limiter
77
+ from sky.security import get_security_guardrails
78
+
79
+ rate_limiter = get_rate_limiter()
80
+ if not rate_limiter.is_allowed("cli_user"):
81
+ console.print("[bold red]Rate limit exceeded. Please wait before making more requests.[/bold red]")
82
+ raise typer.Exit(1)
83
+
84
+ guardrails = get_security_guardrails()
85
+ guardrails.strict_mode = getattr(config, "security_strict_mode", True)
86
+
87
+ is_safe, sanitized_prompt, warning = guardrails.process_user_input(prompt)
88
+ if not is_safe:
89
+ console.print("\n[bold red]🔒 Security Violation Detected[/bold red]")
90
+ console.print("─────────────────────────────────────────")
91
+ if guardrails.violations:
92
+ latest = guardrails.violations[-1]
93
+ console.print(f" [bold]Category:[/bold] {latest[0]}")
94
+ console.print(f" [bold]Pattern:[/bold] {latest[1]}")
95
+ console.print("\n [bold red]Your request was blocked to protect your system.[/bold red]")
96
+ console.print(f" Reason: {warning}")
97
+ console.print("─────────────────────────────────────────")
98
+ console.print(" [dim][Log] Security event recorded in audit log[/dim]\n")
99
+ raise typer.Exit(1)
100
+
101
+ prompt = sanitized_prompt
102
+
103
+ models = load_models_config()
104
+ db = get_db()
105
+ session_id = db.create_session(mode, prompt)
106
+
107
+ router = ModelRouter(models, db, session_id)
108
+ approval_gate = ApprovalGate(config, db, session_id)
109
+
110
+ indexer = None
111
+ if getattr(config, "memory_enabled", False):
112
+ try:
113
+ from sky.memory import get_vector_store, get_indexer
114
+ vector_store = get_vector_store(Path.home() / ".sky" / "lancedb", config.memory_model)
115
+ indexer = get_indexer(config, db, vector_store)
116
+ if getattr(config, "memory_auto_index", False) and not quiet:
117
+ console.print("[dim]Checking repository index...[/dim]")
118
+ indexer.index_directory(Path.cwd())
119
+ except Exception as e:
120
+ if not quiet:
121
+ console.print(f"[bold yellow]Warning: Memory initialization failed: {e}[/bold yellow]")
122
+
123
+ engine = FastLoopEngine(config, db, router, session_id, approval_gate=approval_gate, indexer=indexer)
124
+
125
+ messages = [{"role": "user", "content": prompt}]
126
+ db.append_message(session_id, "user", prompt)
127
+
128
+ if not quiet:
129
+ console.print(f"[bold green]Starting {mode.upper()} mode...[/bold green] (Session: {session_id})")
130
+
131
+ try:
132
+ if not quiet:
133
+ with console.status("[bold cyan]Thinking...[/bold cyan]", spinner="dots") as status:
134
+ async for event in engine.run(messages, mode, inject_context=inject_context, max_turns=max_turns):
135
+ if event["type"] == "model_response":
136
+ if event["message"].get("content"):
137
+ from rich.markdown import Markdown
138
+ status.stop()
139
+ console.print(Markdown(event["message"]["content"]))
140
+ status.start()
141
+ elif event["type"] == "tool_results":
142
+ for res in event["results"]:
143
+ status.stop()
144
+ console.print(f"[dim]Tool {res.get('name', 'unknown')} completed.[/dim]")
145
+ status.start()
146
+ elif event["type"] == "error":
147
+ status.stop()
148
+ console.print(f"[bold red]Error:[/bold red] {event['content']}")
149
+ status.start()
150
+ else:
151
+ # Quiet mode: consume events but don't print unless error or final answer
152
+ async for event in engine.run(messages, mode, inject_context=inject_context, max_turns=max_turns):
153
+ if event["type"] == "final_answer":
154
+ print(event.get("content", ""))
155
+ elif event["type"] == "error":
156
+ print(f"Error: {event['content']}", file=sys.stderr)
157
+
158
+ db.update_session_status(session_id, "completed")
159
+ except SkyError as e:
160
+ if not quiet:
161
+ console.print(f"[bold red]Error ({e.code}):[/bold red] {e.message}")
162
+ if e.suggestion:
163
+ console.print(f"[bold yellow]Suggestion:[/bold yellow] {e.suggestion}")
164
+ else:
165
+ print(f"Error: {e.message}", file=sys.stderr)
166
+ db.update_session_status(session_id, "failed")
167
+ except Exception as e:
168
+ if not quiet:
169
+ console.print(f"[bold red]Execution failed:[/bold red] {e}")
170
+ else:
171
+ print(f"Execution failed: {e}", file=sys.stderr)
172
+ db.update_session_status(session_id, "failed")
173
+
174
+
175
+ from typing import Optional
176
+
177
+ @app.command("chat")
178
+ def chat(
179
+ prompt: Optional[str] = typer.Argument(None, help="Initial message (optional)"),
180
+ model: Optional[str] = typer.Option(None, "--model", help="Override model for chat"),
181
+ verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose output"),
182
+ ):
183
+ """
184
+ Start a conversational chat with Sky.
185
+
186
+ Examples:
187
+ sky chat
188
+ sky chat "Who built you?"
189
+ sky chat "What can you do?"
190
+ sky chat "Help me plan a feature"
191
+
192
+ In chat mode, Sky will:
193
+ - Introduce itself as Sky (not ChatGPT)
194
+ - Share information about its creator when asked
195
+ - Answer questions about its capabilities
196
+ - Suggest the right command for your task
197
+ - Redirect to /ask, /agent, /workflow when appropriate
198
+
199
+ Created by Aaditya A (AI/ML Intern at CoRover.ai)
200
+ """
201
+ from sky.config import load_config, load_models_config
202
+ from sky.storage import get_db
203
+ from sky.core.router import ModelRouter
204
+ from sky.core.chat import ChatEngine
205
+ from dotenv import load_dotenv
206
+
207
+ load_dotenv()
208
+ config = load_config()
209
+ models_config = load_models_config()
210
+
211
+ if verbose:
212
+ config.verbose = True
213
+
214
+ db = get_db()
215
+ session_id = db.create_session("chat", prompt or "conversation")
216
+ router = ModelRouter(models_config, db, session_id)
217
+
218
+ engine = ChatEngine(config, models_config, db, router)
219
+ engine.run(prompt)
220
+
221
+
222
+ @app.command()
223
+ def ask(
224
+ prompt: str = typer.Argument(..., help="Your question"),
225
+ context: bool = typer.Option(False, "--context", help="Inject semantic context from indexed repository"),
226
+ verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose output"),
227
+ quiet: bool = typer.Option(False, "--quiet", "-q", help="Quiet mode (minimal output)")
228
+ ) -> None:
229
+ """
230
+ Ask Sky a question in read-only mode.
231
+
232
+ Examples:
233
+ sky ask "What does FastLoopEngine do?"
234
+ sky ask --context "Explain the approval gate"
235
+ """
236
+ asyncio.run(_run_loop("ask", prompt, inject_context=context, quiet=quiet, verbose=verbose))
237
+
238
+
239
+ @app.command()
240
+ def agent(
241
+ prompt: str = typer.Argument(..., help="Your request"),
242
+ context: bool = typer.Option(False, "--context", help="Inject semantic context from indexed repository"),
243
+ workflow: bool = typer.Option(False, "--workflow", help="Enable workflow visualization"),
244
+ max_turns: int = typer.Option(20, "--max-turns", help="Max conversation turns"),
245
+ verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose output"),
246
+ quiet: bool = typer.Option(False, "--quiet", "-q", help="Quiet mode")
247
+ ) -> None:
248
+ """
249
+ Run Sky in agent mode with full tool access.
250
+
251
+ Examples:
252
+ sky agent "Fix the bug in approval.py"
253
+ sky agent --workflow "Add authentication"
254
+ """
255
+ asyncio.run(_run_loop("agent", prompt, inject_context=context, quiet=quiet, verbose=verbose, max_turns=max_turns))
256
+
257
+
258
+ @app.command()
259
+ def plan(prompt: str) -> None:
260
+ """Planning mode to produce a structured plan."""
261
+ asyncio.run(_run_loop("plan", prompt))
262
+
263
+
264
+ @app.command()
265
+ def stats() -> None:
266
+ """Show database statistics and usage."""
267
+ from sky.storage import get_db
268
+ db = get_db()
269
+ stats = db.get_total_usage()
270
+ console.print("[bold blue]Stats[/bold blue]")
271
+ console.print(f"Total sessions: {stats['total_sessions']}")
272
+ console.print(f"Total cost: ${stats['total_cost']:.4f}")
273
+
274
+
275
+ @app.command()
276
+ def audit(session_id: str) -> None:
277
+ """View session audit logs."""
278
+ from sky.storage import get_db
279
+ db = get_db()
280
+ logs = list(db.get_audit_log(session_id))
281
+ console.print(f"[bold cyan]Audit Log for {session_id}[/bold cyan]")
282
+ for log in logs:
283
+ console.print(log)
284
+
285
+
286
+ @app.command()
287
+ def sessions() -> None:
288
+ """List recent sessions."""
289
+ from sky.storage import get_db
290
+ db = get_db()
291
+ sessions = db.list_sessions(limit=10)
292
+ for s in sessions:
293
+ console.print(f"- {s.id} | {s.mode} | {s.status}")
294
+
295
+
296
+ @app.command()
297
+ def resume(session_id: str) -> None:
298
+ """Resume an interrupted session."""
299
+ console.print(f"Resuming {session_id} not fully implemented in stub.")
300
+
301
+
302
+ @app.command("check-tools")
303
+ def check_tools() -> None:
304
+ """List all registered tools available to SKY."""
305
+ from sky.tools.registry import get_tool_schemas, get_tool
306
+ from sky.config import load_config
307
+ from sky.storage import get_db
308
+ from sky.memory import get_vector_store, get_indexer
309
+
310
+ schemas = get_tool_schemas()
311
+ console.print(f"[bold cyan]Registered Tools ({len(schemas)})[/bold cyan]")
312
+ for schema in schemas:
313
+ name = schema["function"]["name"]
314
+ tool_def = get_tool(name)
315
+ risk = tool_def.risk_tier.name if tool_def else "UNKNOWN"
316
+ console.print(f"- [bold]{name}[/bold] (Risk: [yellow]{risk}[/yellow])")
317
+
318
+ config = load_config()
319
+ db = get_db()
320
+ if getattr(config, "memory_enabled", False):
321
+ try:
322
+ vector_store = get_vector_store(Path.home() / ".sky" / "lancedb", config.memory_model)
323
+ indexer = get_indexer(config, db, vector_store)
324
+ stats = indexer.get_stats()
325
+ console.print("\n[bold cyan]Memory Status[/bold cyan]")
326
+ console.print(f"Status: [green]Enabled[/green]")
327
+ console.print(f"Indexed Files: {stats.get('total_files', 0)}")
328
+ console.print(f"Total Chunks: {stats.get('total_chunks', 0)}")
329
+ except Exception as e:
330
+ console.print(f"\n[bold cyan]Memory Status[/bold cyan]")
331
+ console.print(f"Status: [yellow]Error ({e})[/yellow]")
332
+ else:
333
+ console.print("\n[bold cyan]Memory Status[/bold cyan]")
334
+ console.print("Status: [dim]Disabled[/dim]")
335
+
336
+ @app.command("index")
337
+ def index_repo(
338
+ force: bool = typer.Option(False, "--force", help="Force re-index all files"),
339
+ verbose: bool = typer.Option(False, "--verbose", "-v", help="Show progress")
340
+ ) -> None:
341
+ """
342
+ Index repository for semantic search.
343
+
344
+ Examples:
345
+ sky index # Index current repo
346
+ sky index --force # Force re-index all files
347
+ sky index --verbose # Show progress
348
+ """
349
+ from sky.config import load_config
350
+ from sky.storage import get_db
351
+ from sky.memory import VectorStoreManager, RepoIndexer
352
+
353
+ config = load_config()
354
+ db = get_db()
355
+
356
+ if not getattr(config, "memory_enabled", False):
357
+ console.print("[bold red]Error: Memory is disabled in configuration.[/bold red]")
358
+ raise typer.Exit(1)
359
+
360
+ try:
361
+ vector_store = get_vector_store(Path.home() / ".sky" / "lancedb", config.memory_model)
362
+ if force:
363
+ console.print("[yellow]Clearing existing vector store...[/yellow]")
364
+ vector_store.clear()
365
+
366
+ indexer = RepoIndexer(config, db, vector_store)
367
+
368
+ from rich.progress import Progress, SpinnerColumn, TextColumn, BarColumn
369
+ with Progress(
370
+ SpinnerColumn(),
371
+ TextColumn("[progress.description]{task.description}"),
372
+ BarColumn(),
373
+ transient=True,
374
+ ) as progress:
375
+ task = progress.add_task("[cyan]Indexing repository...", total=None)
376
+
377
+ def update_progress(msg: str):
378
+ progress.update(task, description=f"[cyan]Indexing: {msg}")
379
+
380
+ stats = indexer.index_directory(Path.cwd(), progress_callback=update_progress if verbose else None)
381
+
382
+ console.print("[bold green]Indexing Complete![/bold green]")
383
+ console.print(f"Files Indexed: {stats.get('indexed', 0)}")
384
+ console.print(f"Files Skipped (unchanged/binary): {stats.get('skipped', 0)}")
385
+ console.print(f"Errors: {stats.get('errors', 0)}")
386
+ console.print(f"Deleted Files Pruned: {stats.get('pruned', 0)}")
387
+ console.print("[dim]Embeddings generated locally (free)[/dim]")
388
+ except Exception as e:
389
+ console.print(f"[bold red]Indexing failed: {e}[/bold red]")
390
+
391
+ @app.command("context")
392
+ def search_context_cmd(query: str, top_k: int = typer.Option(5, "--top-k", help="Number of results")) -> None:
393
+ """Search the indexed repository."""
394
+ from sky.config import load_config
395
+ from sky.memory import VectorStoreManager
396
+
397
+ config = load_config()
398
+ if not getattr(config, "memory_enabled", False):
399
+ console.print("[bold red]Error: Memory is disabled in configuration.[/bold red]")
400
+ raise typer.Exit(1)
401
+
402
+ try:
403
+ vector_store = get_vector_store(Path.home() / ".sky" / "lancedb", config.memory_model)
404
+ results = vector_store.search(query, top_k=top_k)
405
+
406
+ from rich.table import Table
407
+ if not results:
408
+ console.print("[yellow]No relevant context found.[/yellow]")
409
+ return
410
+
411
+ console.print(f"[bold green]Found {len(results)} context chunks:[/bold green]")
412
+ table = Table(show_header=True, header_style="bold magenta")
413
+ table.add_column("Score", style="cyan", width=6)
414
+ table.add_column("File Path", style="green", width=30)
415
+ table.add_column("Snippet", style="dim")
416
+
417
+ for res in results:
418
+ score = res.get("score", 0.0)
419
+ path = res.get("file_path", "")
420
+ snippet = res.get("content", "").replace("\n", " ")[:60] + "..."
421
+ table.add_row(f"{score:.2f}", path, snippet)
422
+
423
+ console.print(table)
424
+ except Exception as e:
425
+ console.print(f"[bold red]Search failed: {e}[/bold red]")
426
+
427
+ from typing import Optional
428
+
429
+ async def run_workflow_with_display(engine, goal: str, max_retries: int, resume: Optional[str] = None, quiet: bool = False):
430
+ from rich.markdown import Markdown
431
+ from sky.errors import SkyError
432
+
433
+ # State tracking
434
+ result = {}
435
+
436
+ try:
437
+ if quiet:
438
+ async for event in engine.run_streaming(goal, max_retries=max_retries, resume_from=resume):
439
+ if event.get("type") == "complete":
440
+ print(event.get("summary", ""))
441
+ result = event
442
+ elif event.get("type") == "error":
443
+ print(f"Error: {event.get('content')}", file=sys.stderr)
444
+ result = event
445
+ return result
446
+
447
+ with console.status("[bold cyan]Thinking...[/bold cyan]", spinner="dots") as status:
448
+ async for event in engine.run_streaming(goal, max_retries=max_retries, resume_from=resume):
449
+ event_type = event.get("type")
450
+
451
+ if event_type == "step":
452
+ status.stop()
453
+ console.print(f" [cyan]{event['message']}[/cyan]")
454
+ status.start()
455
+
456
+ elif event_type == "step_complete":
457
+ status.stop()
458
+ console.print(f" [green] {event['step']}: Done[/green]")
459
+ if "result" in event:
460
+ res_str = str(event['result'])
461
+ if len(res_str) > 200:
462
+ res_str = res_str[:197] + "..."
463
+ console.print(f" [dim]{res_str}[/dim]")
464
+ status.start()
465
+
466
+ elif event_type == "subagent_start":
467
+ status.stop()
468
+ console.print(f" [bold magenta] {event['role']}: {event['task']}[/bold magenta]")
469
+ status.start()
470
+
471
+ elif event_type == "subagent_complete":
472
+ status.stop()
473
+ console.print(f" [bold green] {event['role']} complete[/bold green]")
474
+ if "summary" in event and event["summary"]:
475
+ res_str = str(event['summary'])
476
+ if len(res_str) > 200:
477
+ res_str = res_str[:197] + "..."
478
+ console.print(f" [dim]{res_str}[/dim]")
479
+ status.start()
480
+
481
+ elif event_type == "tool_call":
482
+ status.stop()
483
+ args_str = str(event.get('args', ''))
484
+ if len(args_str) > 200:
485
+ args_str = args_str[:197] + "..."
486
+ console.print(f" [dim] Calling: {event['tool_name']}({args_str})[/dim]")
487
+ status.start()
488
+
489
+ elif event_type == "tool_result":
490
+ status.stop()
491
+ res_str = str(event.get('result', ''))
492
+ if len(res_str) > 200:
493
+ res_str = res_str[:197] + "..."
494
+ console.print(f" [dim] {event.get('tool_name', 'tool')} completed: {res_str}[/dim]")
495
+ status.start()
496
+
497
+ elif event_type == "error":
498
+ status.stop()
499
+ console.print(f" [bold red]Error:[/bold red] {event.get('content')}")
500
+ status.start()
501
+
502
+ elif event_type == "complete":
503
+ status.stop()
504
+ console.print(f"\n[bold green]Workflow Complete![/bold green]")
505
+ if "summary" in event:
506
+ console.print(Markdown(event["summary"]))
507
+ result = event
508
+ status.start()
509
+ except SkyError as e:
510
+ if not quiet:
511
+ console.print(f"\n[bold red]Error ({e.code}):[/bold red] {e.message}")
512
+ if e.suggestion:
513
+ console.print(f"[bold yellow]Suggestion:[/bold yellow] {e.suggestion}")
514
+ else:
515
+ print(f"Error: {e.message}", file=sys.stderr)
516
+ except Exception as e:
517
+ if not quiet:
518
+ console.print(f"\n[bold red]Workflow crashed:[/bold red] {e}")
519
+ else:
520
+ print(f"Workflow crashed: {e}", file=sys.stderr)
521
+
522
+ return result
523
+
524
+ @app.command("workflow")
525
+ def workflow(
526
+ goal: str = typer.Argument(..., help="The goal to achieve"),
527
+ max_retries: int = typer.Option(3, "--max-retries", help="Maximum retry attempts"),
528
+ resume: Optional[str] = typer.Option(None, "--resume", help="Session ID to resume"),
529
+ verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose output"),
530
+ quiet: bool = typer.Option(False, "--quiet", "-q", help="Quiet mode")
531
+ ) -> None:
532
+ """
533
+ Run a multi-step workflow with subagents.
534
+
535
+ Examples:
536
+ sky workflow "Add a docstring to approval.py"
537
+ sky workflow --resume <session_id> "Continue"
538
+ """
539
+ from sky.config import load_config, load_models_config
540
+ from sky.core.approval import ApprovalGate
541
+ from sky.core.router import ModelRouter
542
+ from sky.storage import get_db
543
+ from dotenv import load_dotenv
544
+ load_dotenv()
545
+
546
+ if not quiet:
547
+ console.print(f"[bold blue]Starting WORKFLOW mode...[/bold blue]")
548
+ console.print(f"Goal: {goal}\n")
549
+
550
+ # Load config
551
+ config = load_config()
552
+ models_config = load_models_config()
553
+ if verbose:
554
+ config.verbose = True
555
+
556
+ # Initialize components
557
+ db = get_db()
558
+ session_id = resume if resume else db.create_session("workflow", goal)
559
+ router = ModelRouter(models_config, db, session_id)
560
+ approval_gate = ApprovalGate(config, db, session_id)
561
+
562
+ # Initialize memory if enabled
563
+ indexer = None
564
+ if config.memory_enabled:
565
+ from sky.memory.vectorstore import get_vector_store
566
+ from sky.memory.indexer import RepoIndexer
567
+ vector_store = get_vector_store(Path.cwd() / ".sky" / "vectors")
568
+ indexer = RepoIndexer(config, db, vector_store)
569
+
570
+ # Initialize workflow engine (Lazy import to preserve < 150ms startup)
571
+ from sky.core.workflow import WorkflowEngine
572
+ engine = WorkflowEngine(
573
+ config=config,
574
+ db=db,
575
+ router=router,
576
+ approval_gate=approval_gate,
577
+ indexer=indexer,
578
+ )
579
+
580
+ import asyncio
581
+ asyncio.run(run_workflow_with_display(engine, goal, max_retries, resume, quiet=quiet))
582
+
583
+
584
+ @app.command("init")
585
+ def init(
586
+ provider: str = typer.Option(None, "--provider", help="Provider: groq, nim, hybrid"),
587
+ global_install: bool = typer.Option(False, "--global", help="Install globally (~/.sky/)"),
588
+ force: bool = typer.Option(False, "--force", help="Overwrite existing config")
589
+ ):
590
+ """
591
+ Interactive setup for Sky.
592
+
593
+ Examples:
594
+ sky init # Interactive setup
595
+ sky init --provider groq # Quick setup with Groq
596
+ """
597
+ from pathlib import Path
598
+
599
+ console.print("[bold blue]Initializing SKY Project...[/bold blue]")
600
+
601
+ # 1. Ask for NIM API Key
602
+ nim_key = typer.prompt("Enter your NVIDIA NIM API Key (or press Enter to skip)", default="", show_default=False)
603
+ groq_key = typer.prompt("Enter your GROQ API Key (or press Enter to skip)", default="", show_default=False)
604
+
605
+ env_content = ""
606
+ env_path = Path(".env")
607
+ if env_path.exists():
608
+ env_content = env_path.read_text()
609
+
610
+ if nim_key and "NIM_API_KEY" not in env_content:
611
+ with open(env_path, "a") as f:
612
+ f.write(f"\nNIM_API_KEY={nim_key}\n")
613
+ console.print("[green]Added NIM_API_KEY to .env[/green]")
614
+
615
+ if groq_key and "GROQ_API_KEY" not in env_content:
616
+ with open(env_path, "a") as f:
617
+ f.write(f"\nGROQ_API_KEY={groq_key}\n")
618
+ console.print("[green]Added GROQ_API_KEY to .env[/green]")
619
+
620
+ # 2. Generate local models.yaml
621
+ models_yaml_path = Path("models.yaml")
622
+ if models_yaml_path.exists():
623
+ if typer.confirm("models.yaml already exists. Overwrite?"):
624
+ _write_default_models_yaml(models_yaml_path)
625
+ console.print("[green]Overwrote models.yaml[/green]")
626
+ else:
627
+ _write_default_models_yaml(models_yaml_path)
628
+ console.print("[green]Created models.yaml[/green]")
629
+
630
+ console.print("\n[bold green]SKY initialization complete![/bold green]")
631
+ console.print("Run [cyan]sky check-providers[/cyan] to verify your setup.")
632
+
633
+ def _write_default_models_yaml(path: Path):
634
+ content = """providers:
635
+ groq:
636
+ base_url: "https://api.groq.com/openai/v1"
637
+ timeout: 30
638
+ models:
639
+ - id: "groq/compound-mini"
640
+ description: "Ultra-fast routing (0.1s)"
641
+ - id: "openai/gpt-oss-120b"
642
+ description: "Best general conversation"
643
+ - id: "meta-models/Muse-Glimmer-30B"
644
+ description: "Dedicated reasoning & planning"
645
+ - id: "qwen/qwen3.6-27b"
646
+ description: "Best-in-class tool calling"
647
+
648
+ nim:
649
+ base_url: "https://integrate.api.nvidia.com/v1"
650
+ timeout: 60
651
+ models:
652
+ - id: "mistralai/devstral-2"
653
+ description: "Purpose-built for agentic coding"
654
+ - id: "nvidia/llama-3.1-nemotron-70b-instruct"
655
+ description: "Reliable backup model"
656
+
657
+ roles:
658
+ general:
659
+ provider: "groq"
660
+ model_id: "openai/gpt-oss-120b"
661
+ temperature: 0.7
662
+ description: "User interaction, chat, explanations"
663
+
664
+ planning:
665
+ provider: "groq"
666
+ model_id: "meta-models/Muse-Glimmer-30B"
667
+ temperature: 0.3
668
+ description: "Task decomposition, structured planning"
669
+
670
+ reviewer:
671
+ provider: "groq"
672
+ model_id: "meta-models/Muse-Glimmer-30B"
673
+ temperature: 0.3
674
+ description: "Code review, quality analysis"
675
+
676
+ routing:
677
+ provider: "groq"
678
+ model_id: "groq/compound-mini"
679
+ temperature: 0.0
680
+ description: "Intent classification, simple decisions"
681
+
682
+ fast_loop:
683
+ provider: "groq"
684
+ model_id: "qwen/qwen3.6-27b"
685
+ temperature: 0.1
686
+ description: "Parallel tool calling, function execution"
687
+
688
+ coder:
689
+ provider: "nim"
690
+ model_id: "mistralai/devstral-2"
691
+ temperature: 0.1
692
+ description: "Agentic coding, code generation"
693
+
694
+ tester:
695
+ provider: "nim"
696
+ model_id: "mistralai/devstral-2"
697
+ temperature: 0.1
698
+ description: "Test generation, pattern recognition"
699
+
700
+ fallback:
701
+ provider: "nim"
702
+ model_id: "nvidia/llama-3.1-nemotron-70b-instruct"
703
+ temperature: 0.1
704
+ description: "Reliable backup when primary fails"
705
+ """
706
+ path.write_text(content)
707
+
708
+
709
+ def validate_model(provider: str, model_id: str) -> bool:
710
+ """Check if the model exists on the provider."""
711
+ import os
712
+ import httpx
713
+
714
+ if provider == "groq":
715
+ key = os.getenv("GROQ_API_KEY")
716
+ if key:
717
+ try:
718
+ res = httpx.get("https://api.groq.com/openai/v1/models", headers={"Authorization": f"Bearer {key}"})
719
+ if res.status_code == 200:
720
+ models = [m["id"] for m in res.json().get("data", [])]
721
+ if model_id not in models:
722
+ console.print(f"[yellow]Warning: Model {model_id} not found in Groq models list![/yellow]")
723
+ except:
724
+ pass
725
+ elif provider == "nim":
726
+ key = os.getenv("NIM_API_KEY")
727
+ if key:
728
+ try:
729
+ res = httpx.get("https://integrate.api.nvidia.com/v1/models", headers={"Authorization": f"Bearer {key}"})
730
+ if res.status_code == 200:
731
+ models = [m["id"] for m in res.json().get("data", [])]
732
+ if model_id not in models:
733
+ console.print(f"[yellow]Warning: Model {model_id} not found in NIM models list![/yellow]")
734
+ except:
735
+ pass
736
+ return True
737
+
738
+ @app.command("check-providers")
739
+ def check_providers():
740
+ """
741
+ Check provider connectivity and configuration.
742
+
743
+ Examples:
744
+ sky check-providers
745
+
746
+ Shows:
747
+ - Status of each provider
748
+ - Available models
749
+ - Model recommendations per role
750
+ """
751
+ import os
752
+ import httpx
753
+ from dotenv import load_dotenv
754
+ from sky.config import load_models_config
755
+ from sky.errors import ProviderError, SkyError
756
+
757
+ load_dotenv()
758
+ console.print("[bold blue]Checking Providers...[/bold blue]\n")
759
+
760
+ models_config = load_models_config()
761
+
762
+ try:
763
+ # Helper to print models
764
+ def print_provider_info(provider_name: str, status_msg: str):
765
+ provider = models_config.providers.get(provider_name)
766
+ console.print(f" Status: {status_msg}")
767
+ if provider and provider.models:
768
+ console.print(" Available Models (configured):")
769
+ for m in provider.models:
770
+ desc = m.description or "No description"
771
+ console.print(f" - {m.id} ({desc})")
772
+
773
+ used_roles = [r for r, c in models_config.roles.items() if c.provider == provider_name]
774
+ if used_roles:
775
+ console.print(f" Used for: {', '.join(used_roles)}")
776
+
777
+ # Check NIM
778
+ console.print("[bold]NVIDIA NIM:[/bold]")
779
+ nim_key = os.getenv("NIM_API_KEY")
780
+ if not nim_key:
781
+ raise ProviderError(
782
+ "NVIDIA_NIM_API_KEY not found",
783
+ code="SKY-001",
784
+ suggestion="Add NVIDIA_NIM_API_KEY to .env or run `sky init`"
785
+ )
786
+ else:
787
+ try:
788
+ nim_base = models_config.providers.get("nim").base_url if models_config.providers.get("nim") else "https://integrate.api.nvidia.com/v1"
789
+ res = httpx.get(f"{nim_base.rstrip('/')}/models", headers={"Authorization": f"Bearer {nim_key}"})
790
+ if res.status_code == 200:
791
+ print_provider_info("nim", "✅ [green]Connected[/green]")
792
+ else:
793
+ print_provider_info("nim", f"❌ [red]Connection failed: {res.status_code} {res.text}[/red]")
794
+ except Exception as e:
795
+ print_provider_info("nim", f"❌ [red]Connection failed: {e}[/red]")
796
+
797
+ console.print()
798
+
799
+ # Check Groq
800
+ console.print("[bold]Groq:[/bold]")
801
+ groq_key = os.getenv("GROQ_API_KEY")
802
+ if not groq_key:
803
+ raise ProviderError(
804
+ "GROQ_API_KEY not found",
805
+ code="SKY-001",
806
+ suggestion="Add GROQ_API_KEY to .env or run `sky init`"
807
+ )
808
+ else:
809
+ try:
810
+ groq_base = models_config.providers.get("groq").base_url if models_config.providers.get("groq") else "https://api.groq.com/openai/v1"
811
+ res = httpx.get(f"{groq_base.rstrip('/')}/models", headers={"Authorization": f"Bearer {groq_key}"})
812
+ if res.status_code == 200:
813
+ print_provider_info("groq", "✅ [green]Connected[/green]")
814
+ else:
815
+ print_provider_info("groq", f"❌ [red]Connection failed: {res.status_code} {res.text}[/red]")
816
+ except Exception as e:
817
+ print_provider_info("groq", f"❌ [red]Connection failed: {e}[/red]")
818
+
819
+ console.print("\n[bold]Model Recommendations:[/bold]")
820
+ console.print(" - [bold]General Interaction:[/bold] openai/gpt-oss-120b (Groq) - Best conversation")
821
+ console.print(" - [bold]Planning/Reviewing:[/bold] Muse Glimmer (Groq) - Best reasoning")
822
+ console.print(" - [bold]Routing:[/bold] groq/compound-mini (Groq) - Fastest (0.1s)")
823
+ console.print(" - [bold]Tool Calling:[/bold] Qwen 27b (Groq) - Best tool use")
824
+ console.print(" - [bold]Coding/Testing:[/bold] Devstral 2 (NIM) - Best SWE-bench (77.6%)")
825
+ except SkyError as e:
826
+ console.print(f"\n[bold red]Error ({e.code}):[/bold red] {e.message}")
827
+ if e.suggestion:
828
+ console.print(f"[bold yellow]Suggestion:[/bold yellow] {e.suggestion}")
829
+ import sys
830
+ sys.exit(1)
831
+
832
+ if __name__ == "__main__":
833
+ app()