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.
@@ -0,0 +1,152 @@
1
+ """Input sanitization and validation."""
2
+
3
+ import re
4
+ import os
5
+ from pathlib import Path
6
+ from typing import Optional, List, Set
7
+ import logging
8
+
9
+ logger = logging.getLogger(__name__)
10
+
11
+ # Security patterns
12
+ COMMAND_INJECTION_PATTERNS = [
13
+ r'[;&|`]', # Shell metacharacters
14
+ r'\$\{.*\}', # Variable substitution
15
+ r'\(\(.*\)\)', # Arithmetic expansion
16
+ r'<\(.*\)', # Process substitution
17
+ r'>\(.*\)', # Process substitution
18
+ ]
19
+
20
+ PATH_TRAVERSAL_PATTERNS = [
21
+ r'\.\./', # Unix path traversal
22
+ r'\.\.\\', # Windows path traversal
23
+ r'\.\.', # Double dot
24
+ ]
25
+
26
+ PROMPT_INJECTION_PATTERNS = [
27
+ r'(?i)ignore (?:the )?(?:previous |above |all )?instructions',
28
+ r'(?i)forget (?:the )?(?:previous |above |all )?instructions',
29
+ r'(?i)disregard (?:the )?(?:previous |above |all )?instructions',
30
+ r'(?i)you are now (?:a|an) ',
31
+ r'(?i)you (?:are|will) (?:be |now )?(?:acting as|pretend|pretending)',
32
+ r'(?i)system (?:prompt|instruction|message)',
33
+ r'(?i)new (?:role|persona|character)',
34
+ r'(?i)override (?:your|the) (?:previous|system)',
35
+ r'(?i)do not (?:follow|obey|listen to)',
36
+ r'<\s*script', # HTML/JS injection
37
+ r'<\s*iframe',
38
+ r'javascript:',
39
+ r'data:',
40
+ r'vbscript:',
41
+ ]
42
+
43
+ def sanitize_input(user_input: str, max_length: int = 10000) -> str:
44
+ """Sanitize user input."""
45
+ if not user_input:
46
+ return ""
47
+
48
+ # Truncate to prevent DoS
49
+ if len(user_input) > max_length:
50
+ user_input = user_input[:max_length]
51
+ logger.warning(f"Input truncated to {max_length} chars")
52
+
53
+ # Remove control characters
54
+ user_input = ''.join(ch for ch in user_input if ord(ch) >= 32 or ch == '\n')
55
+
56
+ # Normalize whitespace
57
+ user_input = re.sub(r'\s+', ' ', user_input)
58
+
59
+ return user_input.strip()
60
+
61
+ def detect_prompt_injection(user_input: str) -> bool:
62
+ """Detect prompt injection attempts."""
63
+ for pattern in PROMPT_INJECTION_PATTERNS:
64
+ if re.search(pattern, user_input):
65
+ logger.warning(f"Prompt injection detected: {pattern}")
66
+ return True
67
+ return False
68
+
69
+ def validate_path(path: str, base_path: Optional[Path] = None) -> bool:
70
+ """Validate path to prevent directory traversal."""
71
+ if not path:
72
+ return False
73
+
74
+ # Check for path traversal patterns
75
+ for pattern in PATH_TRAVERSAL_PATTERNS:
76
+ if re.search(pattern, path):
77
+ logger.warning(f"Path traversal attempt detected: {path}")
78
+ return False
79
+
80
+ # Resolve path
81
+ try:
82
+ resolved = Path(path).resolve()
83
+ if base_path:
84
+ base = Path(base_path).resolve()
85
+ try:
86
+ resolved.relative_to(base)
87
+ except ValueError:
88
+ logger.warning(f"Path outside base directory: {path}")
89
+ return False
90
+ except Exception as e:
91
+ logger.warning(f"Path validation error: {e}")
92
+ return False
93
+
94
+ return True
95
+
96
+ def validate_command(command: str) -> bool:
97
+ """Validate command to prevent injection."""
98
+ if not command:
99
+ return False
100
+
101
+ # Check for command injection patterns
102
+ for pattern in COMMAND_INJECTION_PATTERNS:
103
+ if re.search(pattern, command):
104
+ logger.warning(f"Command injection detected: {pattern}")
105
+ return False
106
+
107
+ return True
108
+
109
+ def validate_file_extension(path: str, allowed_extensions: Optional[Set[str]] = None) -> bool:
110
+ """Validate file extension."""
111
+ if allowed_extensions is None:
112
+ allowed_extensions = {'.py', '.md', '.txt', '.json', '.yaml', '.yml', '.toml', '.sh', '.js', '.ts', '.html', '.css'}
113
+
114
+ ext = Path(path).suffix.lower()
115
+ if ext and ext not in allowed_extensions:
116
+ logger.warning(f"Disallowed file extension: {ext}")
117
+ return False
118
+
119
+ return True
120
+
121
+ def validate_tool_args(tool_name: str, args: dict) -> bool:
122
+ """Validate tool arguments."""
123
+ if tool_name in ['read_file', 'write_file', 'edit_file']:
124
+ path = args.get('path', '')
125
+ if not validate_path(path):
126
+ return False
127
+ if tool_name == 'write_file' and not validate_file_extension(path):
128
+ return False
129
+
130
+ if tool_name == 'bash':
131
+ command = args.get('command', '')
132
+ if not validate_command(command):
133
+ return False
134
+
135
+ if tool_name in ['grep', 'glob']:
136
+ pattern = args.get('pattern', '')
137
+ # Extremely long patterns can cause DoS
138
+ if len(pattern) > 1000:
139
+ logger.warning(f"Excessive pattern length: {len(pattern)}")
140
+ return False
141
+
142
+ return True
143
+
144
+ def validate_context_size(messages: list, max_tokens: int = 8000) -> bool:
145
+ """Validate that message history doesn't exceed token limits."""
146
+ # Rough estimation
147
+ total_chars = sum(len(str(m)) for m in messages)
148
+ estimated_tokens = total_chars / 4 # Rough estimate
149
+ if estimated_tokens > max_tokens:
150
+ logger.warning(f"Context size exceeds limit: {estimated_tokens} tokens")
151
+ return False
152
+ return True
@@ -0,0 +1,21 @@
1
+ """Storage module exports."""
2
+
3
+ from .db import (
4
+ DatabaseManager,
5
+ MessageRecord,
6
+ RepoIndexRecord,
7
+ SessionRecord,
8
+ ToolCallRecord,
9
+ UsageLogRecord,
10
+ get_db,
11
+ )
12
+
13
+ __all__ = [
14
+ "DatabaseManager",
15
+ "SessionRecord",
16
+ "MessageRecord",
17
+ "ToolCallRecord",
18
+ "UsageLogRecord",
19
+ "RepoIndexRecord",
20
+ "get_db",
21
+ ]
sky/storage/db.py ADDED
@@ -0,0 +1,492 @@
1
+ """Database storage manager and record models."""
2
+
3
+ import json
4
+ import sqlite3
5
+ import uuid
6
+ from dataclasses import dataclass
7
+ from datetime import datetime, timezone
8
+ from pathlib import Path
9
+ import logging
10
+ from typing import Any, Dict, Iterator, List, Optional
11
+
12
+ logger = logging.getLogger(__name__)
13
+
14
+
15
+ @dataclass
16
+ class SessionRecord:
17
+ """Record representing an agent session."""
18
+ id: str
19
+ mode: str
20
+ task: str
21
+ status: str
22
+ started_at: str
23
+ ended_at: Optional[str]
24
+ total_cost: float
25
+ total_tokens: int
26
+
27
+
28
+ @dataclass
29
+ class MessageRecord:
30
+ """Record representing a message in a session."""
31
+ id: str
32
+ session_id: str
33
+ role: str
34
+ content: Optional[str]
35
+ tool_call_json: Optional[str]
36
+ created_at: str
37
+
38
+
39
+ @dataclass
40
+ class ToolCallRecord:
41
+ """Record representing a tool call execution."""
42
+ id: str
43
+ session_id: str
44
+ tool_name: str
45
+ args_json: str
46
+ risk_tier: str
47
+ decision: str
48
+ approved_by: Optional[str]
49
+ result_json: Optional[str]
50
+ timestamp: str
51
+
52
+
53
+ @dataclass
54
+ class UsageLogRecord:
55
+ """Record representing LLM usage metrics."""
56
+ id: str
57
+ session_id: str
58
+ model: str
59
+ role_label: str
60
+ prompt_tokens: int
61
+ completion_tokens: int
62
+ cost_estimate: float
63
+ latency_ms: float
64
+ timestamp: str
65
+
66
+
67
+ @dataclass
68
+ class RepoIndexRecord:
69
+ """Record representing a file in the repository index."""
70
+ file_path: str
71
+ content_hash: str
72
+ summary: Optional[str]
73
+ embedding_id: Optional[str]
74
+ last_indexed_at: str
75
+ language: str = "text"
76
+ chunks: int = 0
77
+ last_modified: str = ""
78
+
79
+
80
+ class DatabaseManager:
81
+ """Manages SQLite storage for SKY operations."""
82
+
83
+ def __init__(self, db_path: str, audit_dir: Optional[str] = None) -> None:
84
+ """Initialize database manager."""
85
+ self.db_path = db_path
86
+ self.audit_dir = audit_dir
87
+
88
+ db_dir = Path(db_path).parent
89
+ db_dir.mkdir(parents=True, exist_ok=True)
90
+
91
+ if self.audit_dir:
92
+ Path(self.audit_dir).mkdir(parents=True, exist_ok=True)
93
+
94
+ self._init_db()
95
+
96
+ def _init_db(self) -> None:
97
+ """Create tables and indexes if they don't exist."""
98
+ with sqlite3.connect(self.db_path) as conn:
99
+ conn.execute("PRAGMA journal_mode=WAL")
100
+ cursor = conn.cursor()
101
+
102
+ cursor.execute('''
103
+ CREATE TABLE IF NOT EXISTS sessions (
104
+ id TEXT PRIMARY KEY,
105
+ mode TEXT NOT NULL,
106
+ task TEXT NOT NULL,
107
+ status TEXT NOT NULL,
108
+ started_at TEXT NOT NULL,
109
+ ended_at TEXT,
110
+ total_cost REAL DEFAULT 0.0,
111
+ total_tokens INTEGER DEFAULT 0
112
+ )
113
+ ''')
114
+
115
+ cursor.execute('''
116
+ CREATE TABLE IF NOT EXISTS messages (
117
+ id TEXT PRIMARY KEY,
118
+ session_id TEXT NOT NULL,
119
+ role TEXT NOT NULL,
120
+ content TEXT,
121
+ tool_call_json TEXT,
122
+ created_at TEXT NOT NULL,
123
+ FOREIGN KEY (session_id) REFERENCES sessions (id)
124
+ )
125
+ ''')
126
+ cursor.execute("CREATE INDEX IF NOT EXISTS idx_messages_session ON messages(session_id)")
127
+
128
+ cursor.execute('''
129
+ CREATE TABLE IF NOT EXISTS tool_calls (
130
+ id TEXT PRIMARY KEY,
131
+ session_id TEXT NOT NULL,
132
+ tool_name TEXT NOT NULL,
133
+ args_json TEXT NOT NULL,
134
+ risk_tier TEXT NOT NULL,
135
+ decision TEXT NOT NULL,
136
+ approved_by TEXT,
137
+ result_json TEXT,
138
+ timestamp TEXT NOT NULL,
139
+ FOREIGN KEY (session_id) REFERENCES sessions (id)
140
+ )
141
+ ''')
142
+ cursor.execute("CREATE INDEX IF NOT EXISTS idx_tool_calls_session ON tool_calls(session_id)")
143
+
144
+ cursor.execute('''
145
+ CREATE TABLE IF NOT EXISTS usage_log (
146
+ id TEXT PRIMARY KEY,
147
+ session_id TEXT NOT NULL,
148
+ model TEXT NOT NULL,
149
+ role_label TEXT NOT NULL,
150
+ prompt_tokens INTEGER NOT NULL,
151
+ completion_tokens INTEGER NOT NULL,
152
+ cost_estimate REAL NOT NULL,
153
+ latency_ms REAL NOT NULL,
154
+ timestamp TEXT NOT NULL,
155
+ FOREIGN KEY (session_id) REFERENCES sessions (id)
156
+ )
157
+ ''')
158
+ cursor.execute("CREATE INDEX IF NOT EXISTS idx_usage_session ON usage_log(session_id)")
159
+
160
+ cursor.execute('''
161
+ CREATE TABLE IF NOT EXISTS repo_index (
162
+ file_path TEXT PRIMARY KEY,
163
+ content_hash TEXT NOT NULL,
164
+ summary TEXT,
165
+ embedding_id TEXT,
166
+ last_indexed_at TEXT NOT NULL,
167
+ language TEXT DEFAULT 'text',
168
+ chunks INTEGER DEFAULT 0,
169
+ last_modified TEXT DEFAULT ''
170
+ )
171
+ ''')
172
+
173
+ cursor.execute('''
174
+ CREATE TABLE IF NOT EXISTS repo_metadata (
175
+ key TEXT PRIMARY KEY,
176
+ value TEXT NOT NULL,
177
+ updated_at TEXT NOT NULL
178
+ )
179
+ ''')
180
+
181
+ # Try to add new columns to existing repo_index table if upgrading
182
+ try:
183
+ cursor.execute("ALTER TABLE repo_index ADD COLUMN language TEXT DEFAULT 'text'")
184
+ cursor.execute("ALTER TABLE repo_index ADD COLUMN chunks INTEGER DEFAULT 0")
185
+ cursor.execute("ALTER TABLE repo_index ADD COLUMN last_modified TEXT DEFAULT ''")
186
+ except sqlite3.OperationalError:
187
+ pass # Columns already exist
188
+
189
+ conn.commit()
190
+
191
+ @staticmethod
192
+ def _make_id() -> str:
193
+ """Generate a unique identifier."""
194
+ return str(uuid.uuid4())
195
+
196
+ @staticmethod
197
+ def _now_iso() -> str:
198
+ """Get current UTC timestamp in ISO format."""
199
+ return datetime.now(timezone.utc).isoformat()
200
+
201
+ def _write_audit_log(self, session_id: str, entry: Dict[str, Any]) -> None:
202
+ """Append an entry to the JSONL audit log for a session."""
203
+ if not self.audit_dir:
204
+ return
205
+ log_file = Path(self.audit_dir) / f"{session_id}.jsonl"
206
+ try:
207
+ with open(log_file, "a", encoding="utf-8") as f:
208
+ f.write(json.dumps(entry) + "\n")
209
+ except Exception as e:
210
+ logger.warning(f"Failed to write audit log: {e}")
211
+
212
+ def create_session(self, mode: str, task: str) -> str:
213
+ """Create a new session."""
214
+ session_id = self._make_id()
215
+ now = self._now_iso()
216
+ with sqlite3.connect(self.db_path) as conn:
217
+ conn.execute(
218
+ "INSERT INTO sessions (id, mode, task, status, started_at) VALUES (?, ?, ?, ?, ?)",
219
+ (session_id, mode, task, "active", now)
220
+ )
221
+ self._write_audit_log(session_id, {"event": "session_created", "mode": mode, "task": task, "timestamp": now})
222
+ return session_id
223
+
224
+ def update_session_status(self, session_id: str, status: str) -> None:
225
+ """Update session status."""
226
+ now = self._now_iso()
227
+ ended_at = now if status in ("completed", "failed", "cancelled") else None
228
+ with sqlite3.connect(self.db_path) as conn:
229
+ if ended_at:
230
+ conn.execute(
231
+ "UPDATE sessions SET status = ?, ended_at = ? WHERE id = ?",
232
+ (status, ended_at, session_id)
233
+ )
234
+ else:
235
+ conn.execute("UPDATE sessions SET status = ? WHERE id = ?", (status, session_id))
236
+ self._write_audit_log(session_id, {"event": "session_status_updated", "status": status, "timestamp": now})
237
+
238
+ def update_session_metrics(self, session_id: str, cost: float, tokens: int) -> None:
239
+ """Update session aggregate usage metrics."""
240
+ with sqlite3.connect(self.db_path) as conn:
241
+ conn.execute(
242
+ "UPDATE sessions SET total_cost = total_cost + ?, total_tokens = total_tokens + ? WHERE id = ?",
243
+ (cost, tokens, session_id)
244
+ )
245
+
246
+ def get_session(self, session_id: str) -> Optional[SessionRecord]:
247
+ """Get details for a specific session."""
248
+ with sqlite3.connect(self.db_path) as conn:
249
+ conn.row_factory = sqlite3.Row
250
+ row = conn.execute("SELECT * FROM sessions WHERE id = ?", (session_id,)).fetchone()
251
+ if row:
252
+ return SessionRecord(**dict(row))
253
+ return None
254
+
255
+ def list_sessions(self, limit: int = 20) -> List[SessionRecord]:
256
+ """List the most recent sessions."""
257
+ with sqlite3.connect(self.db_path) as conn:
258
+ conn.row_factory = sqlite3.Row
259
+ rows = conn.execute("SELECT * FROM sessions ORDER BY started_at DESC LIMIT ?", (limit,)).fetchall()
260
+ return [SessionRecord(**dict(row)) for row in rows]
261
+
262
+ def append_message(
263
+ self, session_id: str, role: str, content: Optional[str] = None, tool_calls: Optional[List[Dict[str, Any]]] = None
264
+ ) -> str:
265
+ """Append a chat message to a session."""
266
+ msg_id = self._make_id()
267
+ now = self._now_iso()
268
+ tool_call_json = json.dumps(tool_calls) if tool_calls else None
269
+
270
+ with sqlite3.connect(self.db_path) as conn:
271
+ conn.execute(
272
+ "INSERT INTO messages (id, session_id, role, content, tool_call_json, created_at) VALUES (?, ?, ?, ?, ?, ?)",
273
+ (msg_id, session_id, role, content, tool_call_json, now)
274
+ )
275
+
276
+ self._write_audit_log(session_id, {
277
+ "event": "message_appended",
278
+ "message_id": msg_id,
279
+ "role": role,
280
+ "timestamp": now
281
+ })
282
+ return msg_id
283
+
284
+ def get_messages(self, session_id: str) -> List[MessageRecord]:
285
+ """Get all messages for a given session."""
286
+ with sqlite3.connect(self.db_path) as conn:
287
+ conn.row_factory = sqlite3.Row
288
+ rows = conn.execute("SELECT * FROM messages WHERE session_id = ? ORDER BY created_at ASC", (session_id,)).fetchall()
289
+ return [MessageRecord(**dict(row)) for row in rows]
290
+
291
+ def log_tool_call(
292
+ self,
293
+ session_id: str,
294
+ tool_name: str,
295
+ args: Dict[str, Any],
296
+ risk_tier: str,
297
+ decision: str,
298
+ approved_by: Optional[str] = None,
299
+ result: Optional[Any] = None
300
+ ) -> str:
301
+ """Log the execution of a tool."""
302
+ tc_id = self._make_id()
303
+ now = self._now_iso()
304
+ args_json = json.dumps(args)
305
+ result_json = json.dumps(result) if result is not None else None
306
+
307
+ with sqlite3.connect(self.db_path) as conn:
308
+ conn.execute(
309
+ """INSERT INTO tool_calls
310
+ (id, session_id, tool_name, args_json, risk_tier, decision, approved_by, result_json, timestamp)
311
+ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
312
+ (tc_id, session_id, tool_name, args_json, risk_tier, decision, approved_by, result_json, now)
313
+ )
314
+
315
+ self._write_audit_log(session_id, {
316
+ "event": "tool_call",
317
+ "tool_call_id": tc_id,
318
+ "tool_name": tool_name,
319
+ "decision": decision,
320
+ "timestamp": now
321
+ })
322
+ return tc_id
323
+
324
+ def log_usage(
325
+ self,
326
+ session_id: str,
327
+ model: str,
328
+ role_label: str,
329
+ prompt_tokens: int,
330
+ completion_tokens: int,
331
+ cost_estimate: float,
332
+ latency_ms: float
333
+ ) -> str:
334
+ """Log language model usage metrics."""
335
+ usage_id = self._make_id()
336
+ now = self._now_iso()
337
+
338
+ with sqlite3.connect(self.db_path) as conn:
339
+ conn.execute(
340
+ """INSERT INTO usage_log
341
+ (id, session_id, model, role_label, prompt_tokens, completion_tokens, cost_estimate, latency_ms, timestamp)
342
+ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
343
+ (usage_id, session_id, model, role_label, prompt_tokens, completion_tokens, cost_estimate, latency_ms, now)
344
+ )
345
+
346
+ self.update_session_metrics(session_id, cost_estimate, prompt_tokens + completion_tokens)
347
+ return usage_id
348
+
349
+ def get_session_usage(self, session_id: str) -> List[UsageLogRecord]:
350
+ """Get all usage logs for a session."""
351
+ with sqlite3.connect(self.db_path) as conn:
352
+ conn.row_factory = sqlite3.Row
353
+ rows = conn.execute("SELECT * FROM usage_log WHERE session_id = ? ORDER BY timestamp ASC", (session_id,)).fetchall()
354
+ return [UsageLogRecord(**dict(row)) for row in rows]
355
+
356
+ def get_total_usage(self) -> Dict[str, Any]:
357
+ """Get aggregate system usage metrics."""
358
+ with sqlite3.connect(self.db_path) as conn:
359
+ conn.row_factory = sqlite3.Row
360
+ row = conn.execute("""
361
+ SELECT
362
+ COUNT(DISTINCT session_id) as total_sessions,
363
+ SUM(prompt_tokens) as total_prompt_tokens,
364
+ SUM(completion_tokens) as total_completion_tokens,
365
+ SUM(cost_estimate) as total_cost
366
+ FROM usage_log
367
+ """).fetchone()
368
+
369
+ return {
370
+ "total_sessions": row["total_sessions"] or 0,
371
+ "total_prompt_tokens": row["total_prompt_tokens"] or 0,
372
+ "total_completion_tokens": row["total_completion_tokens"] or 0,
373
+ "total_cost": row["total_cost"] or 0.0
374
+ }
375
+
376
+ def get_indexed_files(self) -> List[str]:
377
+ """Get a list of all file paths in the repository index."""
378
+ with sqlite3.connect(self.db_path) as conn:
379
+ rows = conn.execute("SELECT file_path FROM repo_index").fetchall()
380
+ return [row[0] for row in rows]
381
+
382
+ def list_indexed_files(self) -> List[str]:
383
+ """Alias for get_indexed_files."""
384
+ return self.get_indexed_files()
385
+
386
+ def get_index_entry(self, file_path: str) -> Optional[RepoIndexRecord]:
387
+ """Get the repository index entry for a specific file."""
388
+ with sqlite3.connect(self.db_path) as conn:
389
+ conn.row_factory = sqlite3.Row
390
+ row = conn.execute("SELECT * FROM repo_index WHERE file_path = ?", (file_path,)).fetchone()
391
+ if row:
392
+ return RepoIndexRecord(**dict(row))
393
+ return None
394
+
395
+ def update_index_entry(
396
+ self,
397
+ file_path: str,
398
+ content_hash: str,
399
+ summary: Optional[str] = None,
400
+ embedding_id: Optional[str] = None,
401
+ language: str = "text",
402
+ chunks: int = 0,
403
+ last_modified: str = ""
404
+ ) -> None:
405
+ """Update or insert an entry in the repository index."""
406
+ now = self._now_iso()
407
+ with sqlite3.connect(self.db_path) as conn:
408
+ conn.execute(
409
+ """INSERT INTO repo_index (file_path, content_hash, summary, embedding_id, last_indexed_at, language, chunks, last_modified)
410
+ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
411
+ ON CONFLICT(file_path) DO UPDATE SET
412
+ content_hash=excluded.content_hash,
413
+ summary=excluded.summary,
414
+ embedding_id=excluded.embedding_id,
415
+ last_indexed_at=excluded.last_indexed_at,
416
+ language=excluded.language,
417
+ chunks=excluded.chunks,
418
+ last_modified=excluded.last_modified""",
419
+ (file_path, content_hash, summary, embedding_id, now, language, chunks, last_modified)
420
+ )
421
+
422
+ def get_repo_metadata(self, key: str) -> Optional[str]:
423
+ """Get a metadata value by key."""
424
+ with sqlite3.connect(self.db_path) as conn:
425
+ row = conn.execute("SELECT value FROM repo_metadata WHERE key = ?", (key,)).fetchone()
426
+ if row:
427
+ return row[0]
428
+ return None
429
+
430
+ def set_repo_metadata(self, key: str, value: str) -> None:
431
+ """Set a metadata value."""
432
+ now = self._now_iso()
433
+ with sqlite3.connect(self.db_path) as conn:
434
+ conn.execute(
435
+ """INSERT INTO repo_metadata (key, value, updated_at) VALUES (?, ?, ?)
436
+ ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at""",
437
+ (key, value, now)
438
+ )
439
+
440
+ def get_index_stats(self) -> Dict[str, Any]:
441
+ """Get summary stats of the repo index."""
442
+ with sqlite3.connect(self.db_path) as conn:
443
+ conn.row_factory = sqlite3.Row
444
+ row = conn.execute("SELECT COUNT(file_path) as total_files, SUM(chunks) as total_chunks FROM repo_index").fetchone()
445
+ last_idx = conn.execute("SELECT value FROM repo_metadata WHERE key = 'last_indexed'").fetchone()
446
+ return {
447
+ "total_files": row["total_files"] or 0,
448
+ "total_chunks": row["total_chunks"] or 0,
449
+ "last_indexed": last_idx[0] if last_idx else None
450
+ }
451
+
452
+ def remove_index_entry(self, file_path: str) -> None:
453
+ """Remove a file from the repository index."""
454
+ with sqlite3.connect(self.db_path) as conn:
455
+ conn.execute("DELETE FROM repo_index WHERE file_path = ?", (file_path,))
456
+
457
+ def get_audit_log(self, session_id: str) -> Iterator[Dict[str, Any]]:
458
+ """Yield audit log entries for a given session."""
459
+ if not self.audit_dir:
460
+ return
461
+
462
+ log_file = Path(self.audit_dir) / f"{session_id}.jsonl"
463
+ if not log_file.exists():
464
+ return
465
+
466
+ with open(log_file, "r", encoding="utf-8") as f:
467
+ for line in f:
468
+ if line.strip():
469
+ yield json.loads(line)
470
+
471
+ def get_audit_summary(self, session_id: str) -> Dict[str, int]:
472
+ """Get a summary count of audit events for a session."""
473
+ summary: Dict[str, int] = {}
474
+ for entry in self.get_audit_log(session_id):
475
+ event = entry.get("event", "unknown")
476
+ summary[event] = summary.get(event, 0) + 1
477
+ return summary
478
+
479
+
480
+ _db_instance: Optional[DatabaseManager] = None
481
+
482
+
483
+ def get_db(db_path: Optional[str] = None) -> DatabaseManager:
484
+ """Get the global database instance, creating it if necessary."""
485
+ global _db_instance
486
+ if _db_instance is None:
487
+ if db_path is None:
488
+ # Default location
489
+ db_path = str(Path.home() / ".sky" / "storage.db")
490
+ audit_dir = str(Path(db_path).parent / "audit_logs")
491
+ _db_instance = DatabaseManager(db_path, audit_dir)
492
+ return _db_instance
sky/tools/__init__.py ADDED
@@ -0,0 +1,26 @@
1
+ """Tool registry exports."""
2
+
3
+ from .registry import (
4
+ RiskTier,
5
+ ToolDefinition,
6
+ ToolRegistry,
7
+ get_tool,
8
+ get_tool_schemas,
9
+ get_tools_by_risk,
10
+ list_tools,
11
+ register_tool,
12
+ )
13
+
14
+ __all__ = [
15
+ "RiskTier",
16
+ "ToolDefinition",
17
+ "ToolRegistry",
18
+ "register_tool",
19
+ "get_tool",
20
+ "list_tools",
21
+ "get_tools_by_risk",
22
+ "get_tool_schemas",
23
+ ]
24
+
25
+ # Ensure all tool modules are imported so their @register_tool decorators execute
26
+ from . import fs_tools, git_tools, search_tools, shell_tools