-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathdb.go
More file actions
291 lines (241 loc) · 8.83 KB
/
Copy pathdb.go
File metadata and controls
291 lines (241 loc) · 8.83 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
// db.go
package main
import (
"database/sql"
"encoding/json"
"fmt"
"time"
_ "modernc.org/sqlite"
)
// PatternState holds agent behavior tracking data for a session.
// Used for computing tool streaks, retries, and turn depth metrics.
type PatternState struct {
TurnCount int // Number of request-response cycles in this session
LastToolName string // First tool name from previous turn (for streak detection)
ToolStreak int // Consecutive turns where first tool is the same
RetryCount int // Consecutive retry attempts (same tool after error)
SessionToolCount int // Total tool calls in session so far
LastWasError bool // Previous turn's tool resulted in error
PendingToolIDs map[string]string // tool_use_id -> tool_name for result matching
}
type SessionDB struct {
db *sql.DB
}
func NewSessionDB(path string) (*SessionDB, error) {
db, err := sql.Open("sqlite", path)
if err != nil {
return nil, fmt.Errorf("failed to open database: %w", err)
}
// Create tables
schema := `
CREATE TABLE IF NOT EXISTS sessions (
id TEXT PRIMARY KEY,
provider TEXT NOT NULL,
upstream TEXT NOT NULL,
created_at TEXT NOT NULL,
last_activity TEXT NOT NULL,
last_seq INTEGER NOT NULL DEFAULT 0,
last_fingerprint TEXT NOT NULL DEFAULT '',
file_path TEXT NOT NULL,
client_session_id TEXT
);
CREATE TABLE IF NOT EXISTS fingerprints (
fingerprint TEXT PRIMARY KEY,
session_id TEXT NOT NULL,
seq INTEGER NOT NULL,
FOREIGN KEY (session_id) REFERENCES sessions(id)
);
CREATE INDEX IF NOT EXISTS idx_fingerprints_session ON fingerprints(session_id);
CREATE INDEX IF NOT EXISTS idx_sessions_provider ON sessions(provider);
CREATE INDEX IF NOT EXISTS idx_sessions_client_id ON sessions(client_session_id);
`
if _, err := db.Exec(schema); err != nil {
db.Close()
return nil, fmt.Errorf("failed to create schema: %w", err)
}
// Migrations: add columns if they don't exist (ignore "duplicate column" errors)
migrations := []string{
"ALTER TABLE sessions ADD COLUMN client_session_id TEXT",
"ALTER TABLE sessions ADD COLUMN turn_count INTEGER NOT NULL DEFAULT 0",
"ALTER TABLE sessions ADD COLUMN last_tool_name TEXT NOT NULL DEFAULT ''",
"ALTER TABLE sessions ADD COLUMN tool_streak INTEGER NOT NULL DEFAULT 0",
"ALTER TABLE sessions ADD COLUMN retry_count INTEGER NOT NULL DEFAULT 0",
"ALTER TABLE sessions ADD COLUMN session_tool_count INTEGER NOT NULL DEFAULT 0",
"ALTER TABLE sessions ADD COLUMN last_was_error INTEGER NOT NULL DEFAULT 0",
"ALTER TABLE sessions ADD COLUMN pending_tool_ids TEXT NOT NULL DEFAULT '{}'",
}
for _, migration := range migrations {
db.Exec(migration) // Ignore errors - column may already exist
}
// Create index for client_session_id (may already exist from schema)
db.Exec("CREATE INDEX IF NOT EXISTS idx_sessions_client_id ON sessions(client_session_id)")
return &SessionDB{db: db}, nil
}
func (s *SessionDB) Close() error {
return s.db.Close()
}
func (s *SessionDB) CreateSession(id, provider, upstream, filePath string) error {
now := time.Now().UTC().Format(time.RFC3339)
_, err := s.db.Exec(`
INSERT INTO sessions (id, provider, upstream, created_at, last_activity, file_path, last_seq)
VALUES (?, ?, ?, ?, ?, ?, 1)
`, id, provider, upstream, now, now, filePath)
return err
}
// UpdateSessionSeq updates the session's last sequence number
func (s *SessionDB) UpdateSessionSeq(sessionID string, seq int) error {
now := time.Now().UTC().Format(time.RFC3339)
_, err := s.db.Exec(`
UPDATE sessions
SET last_activity = ?, last_seq = ?
WHERE id = ?
`, now, seq, sessionID)
return err
}
func (s *SessionDB) UpdateSessionFingerprint(sessionID string, seq int, fingerprint string) error {
now := time.Now().UTC().Format(time.RFC3339)
// Update session
_, err := s.db.Exec(`
UPDATE sessions
SET last_activity = ?, last_seq = ?, last_fingerprint = ?
WHERE id = ?
`, now, seq, fingerprint, sessionID)
if err != nil {
return err
}
// Insert fingerprint mapping
_, err = s.db.Exec(`
INSERT OR REPLACE INTO fingerprints (fingerprint, session_id, seq)
VALUES (?, ?, ?)
`, fingerprint, sessionID, seq)
return err
}
func (s *SessionDB) FindByFingerprint(fingerprint string) (sessionID string, seq int, err error) {
row := s.db.QueryRow(`
SELECT session_id, seq FROM fingerprints WHERE fingerprint = ?
`, fingerprint)
err = row.Scan(&sessionID, &seq)
if err == sql.ErrNoRows {
return "", 0, nil
}
return sessionID, seq, err
}
func (s *SessionDB) GetLatestFingerprint(sessionID string) (fingerprint string, seq int, err error) {
row := s.db.QueryRow(`
SELECT last_fingerprint, last_seq FROM sessions WHERE id = ?
`, sessionID)
err = row.Scan(&fingerprint, &seq)
if err == sql.ErrNoRows {
return "", 0, nil
}
return fingerprint, seq, err
}
func (s *SessionDB) GetSession(id string) (provider, upstream, filePath string, err error) {
row := s.db.QueryRow(`
SELECT provider, upstream, file_path FROM sessions WHERE id = ?
`, id)
err = row.Scan(&provider, &upstream, &filePath)
return
}
// CreateSessionWithClientID creates a new session with a client-provided session ID
func (s *SessionDB) CreateSessionWithClientID(id, clientSessionID, provider, upstream, filePath string) error {
now := time.Now().UTC().Format(time.RFC3339)
_, err := s.db.Exec(`
INSERT INTO sessions (id, client_session_id, provider, upstream, created_at, last_activity, file_path, last_seq)
VALUES (?, ?, ?, ?, ?, ?, ?, 1)
`, id, clientSessionID, provider, upstream, now, now, filePath)
return err
}
// FindByClientSessionID finds a session by its client-provided session ID
func (s *SessionDB) FindByClientSessionID(clientSessionID string) (sessionID string, err error) {
row := s.db.QueryRow(`
SELECT id FROM sessions WHERE client_session_id = ?
`, clientSessionID)
err = row.Scan(&sessionID)
if err == sql.ErrNoRows {
return "", nil
}
return sessionID, err
}
// GetSessionWithClientID gets a session including its client session ID and last sequence
func (s *SessionDB) GetSessionWithClientID(id string) (provider, upstream, filePath string, lastSeq int, err error) {
row := s.db.QueryRow(`
SELECT provider, upstream, file_path, last_seq FROM sessions WHERE id = ?
`, id)
err = row.Scan(&provider, &upstream, &filePath, &lastSeq)
return
}
// LoadPatternState loads pattern tracking state for a session.
// Returns nil with no error if session doesn't exist.
func (s *SessionDB) LoadPatternState(sessionID string) (*PatternState, error) {
row := s.db.QueryRow(`
SELECT turn_count, last_tool_name, tool_streak, retry_count,
session_tool_count, last_was_error, pending_tool_ids
FROM sessions WHERE id = ?
`, sessionID)
var turnCount, toolStreak, retryCount, sessionToolCount int
var lastToolName, pendingToolIDsJSON string
var lastWasError int
err := row.Scan(&turnCount, &lastToolName, &toolStreak, &retryCount,
&sessionToolCount, &lastWasError, &pendingToolIDsJSON)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
pendingToolIDs := make(map[string]string)
if pendingToolIDsJSON != "" && pendingToolIDsJSON != "{}" {
if err := json.Unmarshal([]byte(pendingToolIDsJSON), &pendingToolIDs); err != nil {
return nil, fmt.Errorf("failed to unmarshal pending_tool_ids: %w", err)
}
}
return &PatternState{
TurnCount: turnCount,
LastToolName: lastToolName,
ToolStreak: toolStreak,
RetryCount: retryCount,
SessionToolCount: sessionToolCount,
LastWasError: lastWasError != 0,
PendingToolIDs: pendingToolIDs,
}, nil
}
// UpdatePatternState persists pattern tracking state for a session.
func (s *SessionDB) UpdatePatternState(sessionID string, state *PatternState) error {
pendingToolIDsJSON, err := json.Marshal(state.PendingToolIDs)
if err != nil {
return fmt.Errorf("failed to marshal pending_tool_ids: %w", err)
}
lastWasError := 0
if state.LastWasError {
lastWasError = 1
}
_, err = s.db.Exec(`
UPDATE sessions
SET turn_count = ?, last_tool_name = ?, tool_streak = ?, retry_count = ?,
session_tool_count = ?, last_was_error = ?, pending_tool_ids = ?
WHERE id = ?
`, state.TurnCount, state.LastToolName, state.ToolStreak, state.RetryCount,
state.SessionToolCount, lastWasError, string(pendingToolIDsJSON), sessionID)
return err
}
// ClearMatchedToolID removes a tool ID from pending_tool_ids and returns the tool name.
// Returns empty string if the tool ID was not found.
func (s *SessionDB) ClearMatchedToolID(sessionID, toolUseID string) (string, error) {
state, err := s.LoadPatternState(sessionID)
if err != nil {
return "", err
}
if state == nil {
return "", nil
}
toolName, exists := state.PendingToolIDs[toolUseID]
if !exists {
return "", nil
}
delete(state.PendingToolIDs, toolUseID)
if err := s.UpdatePatternState(sessionID, state); err != nil {
return "", err
}
return toolName, nil
}