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/security/sanitize.py
ADDED
|
@@ -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
|
sky/storage/__init__.py
ADDED
|
@@ -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
|