first commit

This commit is contained in:
2026-03-08 15:40:34 +07:00
commit 8dc496b626
159 changed files with 27932 additions and 0 deletions
+31
View File
@@ -0,0 +1,31 @@
// Code generated by sqlc. DO NOT EDIT.
// versions:
// sqlc v1.30.0
package db
import (
"context"
"database/sql"
)
type DBTX interface {
ExecContext(context.Context, string, ...interface{}) (sql.Result, error)
PrepareContext(context.Context, string) (*sql.Stmt, error)
QueryContext(context.Context, string, ...interface{}) (*sql.Rows, error)
QueryRowContext(context.Context, string, ...interface{}) *sql.Row
}
func New(db DBTX) *Queries {
return &Queries{db: db}
}
type Queries struct {
db DBTX
}
func (q *Queries) WithTx(tx *sql.Tx) *Queries {
return &Queries{
db: tx,
}
}
+58
View File
@@ -0,0 +1,58 @@
-- Sessions table: replaces noted CLI dependency for session persistence.
CREATE TABLE IF NOT EXISTS sessions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
title TEXT NOT NULL DEFAULT '',
model TEXT NOT NULL DEFAULT '',
mode TEXT NOT NULL DEFAULT 'BUILD',
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now')),
updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
-- Session messages: individual chat entries within a session.
CREATE TABLE IF NOT EXISTS session_messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id INTEGER NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
role TEXT NOT NULL, -- 'user', 'assistant', 'tool', 'system', 'error'
content TEXT NOT NULL DEFAULT '',
tool_name TEXT NOT NULL DEFAULT '',
tool_args TEXT NOT NULL DEFAULT '',
is_error INTEGER NOT NULL DEFAULT 0,
thinking TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
CREATE INDEX IF NOT EXISTS idx_session_messages_session_id ON session_messages(session_id);
-- Tool permissions: per-tool allow/deny/always-allow.
CREATE TABLE IF NOT EXISTS tool_permissions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
tool_name TEXT NOT NULL UNIQUE,
policy TEXT NOT NULL DEFAULT 'ask', -- 'allow', 'deny', 'ask'
updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
-- Token usage stats: per-turn tracking.
CREATE TABLE IF NOT EXISTS token_stats (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id INTEGER NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
turn INTEGER NOT NULL DEFAULT 0,
eval_count INTEGER NOT NULL DEFAULT 0,
prompt_tokens INTEGER NOT NULL DEFAULT 0,
model TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
CREATE INDEX IF NOT EXISTS idx_token_stats_session_id ON token_stats(session_id);
-- File changes: files modified by agent during a session.
CREATE TABLE IF NOT EXISTS file_changes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id INTEGER NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
file_path TEXT NOT NULL,
tool_name TEXT NOT NULL DEFAULT '',
added INTEGER NOT NULL DEFAULT 0,
removed INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
CREATE INDEX IF NOT EXISTS idx_file_changes_session_id ON file_changes(session_id);
+53
View File
@@ -0,0 +1,53 @@
// Code generated by sqlc. DO NOT EDIT.
// versions:
// sqlc v1.30.0
package db
type FileChange struct {
ID int64 `json:"id"`
SessionID int64 `json:"session_id"`
FilePath string `json:"file_path"`
ToolName string `json:"tool_name"`
Added int64 `json:"added"`
Removed int64 `json:"removed"`
CreatedAt string `json:"created_at"`
}
type Session struct {
ID int64 `json:"id"`
Title string `json:"title"`
Model string `json:"model"`
Mode string `json:"mode"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
type SessionMessage struct {
ID int64 `json:"id"`
SessionID int64 `json:"session_id"`
Role string `json:"role"`
Content string `json:"content"`
ToolName string `json:"tool_name"`
ToolArgs string `json:"tool_args"`
IsError int64 `json:"is_error"`
Thinking string `json:"thinking"`
CreatedAt string `json:"created_at"`
}
type TokenStat struct {
ID int64 `json:"id"`
SessionID int64 `json:"session_id"`
Turn int64 `json:"turn"`
EvalCount int64 `json:"eval_count"`
PromptTokens int64 `json:"prompt_tokens"`
Model string `json:"model"`
CreatedAt string `json:"created_at"`
}
type ToolPermission struct {
ID int64 `json:"id"`
ToolName string `json:"tool_name"`
Policy string `json:"policy"`
UpdatedAt string `json:"updated_at"`
}
+100
View File
@@ -0,0 +1,100 @@
// Code generated by sqlc. DO NOT EDIT.
// versions:
// sqlc v1.30.0
// source: permissions.sql
package db
import (
"context"
)
const deleteToolPermission = `-- name: DeleteToolPermission :exec
DELETE FROM tool_permissions WHERE tool_name = ?
`
func (q *Queries) DeleteToolPermission(ctx context.Context, toolName string) error {
_, err := q.db.ExecContext(ctx, deleteToolPermission, toolName)
return err
}
const getToolPermission = `-- name: GetToolPermission :one
SELECT id, tool_name, policy, updated_at FROM tool_permissions WHERE tool_name = ?
`
func (q *Queries) GetToolPermission(ctx context.Context, toolName string) (ToolPermission, error) {
row := q.db.QueryRowContext(ctx, getToolPermission, toolName)
var i ToolPermission
err := row.Scan(
&i.ID,
&i.ToolName,
&i.Policy,
&i.UpdatedAt,
)
return i, err
}
const listToolPermissions = `-- name: ListToolPermissions :many
SELECT id, tool_name, policy, updated_at FROM tool_permissions ORDER BY tool_name ASC
`
func (q *Queries) ListToolPermissions(ctx context.Context) ([]ToolPermission, error) {
rows, err := q.db.QueryContext(ctx, listToolPermissions)
if err != nil {
return nil, err
}
defer rows.Close()
items := []ToolPermission{}
for rows.Next() {
var i ToolPermission
if err := rows.Scan(
&i.ID,
&i.ToolName,
&i.Policy,
&i.UpdatedAt,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const resetToolPermissions = `-- name: ResetToolPermissions :exec
DELETE FROM tool_permissions
`
func (q *Queries) ResetToolPermissions(ctx context.Context) error {
_, err := q.db.ExecContext(ctx, resetToolPermissions)
return err
}
const upsertToolPermission = `-- name: UpsertToolPermission :one
INSERT INTO tool_permissions (tool_name, policy)
VALUES (?, ?)
ON CONFLICT(tool_name) DO UPDATE SET policy = excluded.policy, updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now')
RETURNING id, tool_name, policy, updated_at
`
type UpsertToolPermissionParams struct {
ToolName string `json:"tool_name"`
Policy string `json:"policy"`
}
func (q *Queries) UpsertToolPermission(ctx context.Context, arg UpsertToolPermissionParams) (ToolPermission, error) {
row := q.db.QueryRowContext(ctx, upsertToolPermission, arg.ToolName, arg.Policy)
var i ToolPermission
err := row.Scan(
&i.ID,
&i.ToolName,
&i.Policy,
&i.UpdatedAt,
)
return i, err
}
+17
View File
@@ -0,0 +1,17 @@
-- name: GetToolPermission :one
SELECT * FROM tool_permissions WHERE tool_name = ?;
-- name: UpsertToolPermission :one
INSERT INTO tool_permissions (tool_name, policy)
VALUES (?, ?)
ON CONFLICT(tool_name) DO UPDATE SET policy = excluded.policy, updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now')
RETURNING *;
-- name: ListToolPermissions :many
SELECT * FROM tool_permissions ORDER BY tool_name ASC;
-- name: DeleteToolPermission :exec
DELETE FROM tool_permissions WHERE tool_name = ?;
-- name: ResetToolPermissions :exec
DELETE FROM tool_permissions;
+27
View File
@@ -0,0 +1,27 @@
-- name: CreateSession :one
INSERT INTO sessions (title, model, mode) VALUES (?, ?, ?) RETURNING *;
-- name: GetSession :one
SELECT * FROM sessions WHERE id = ?;
-- name: ListSessions :many
SELECT * FROM sessions ORDER BY updated_at DESC LIMIT ?;
-- name: UpdateSessionTitle :exec
UPDATE sessions SET title = ?, updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now') WHERE id = ?;
-- name: UpdateSessionTimestamp :exec
UPDATE sessions SET updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now') WHERE id = ?;
-- name: DeleteSession :exec
DELETE FROM sessions WHERE id = ?;
-- name: CreateSessionMessage :one
INSERT INTO session_messages (session_id, role, content, tool_name, tool_args, is_error, thinking)
VALUES (?, ?, ?, ?, ?, ?, ?) RETURNING *;
-- name: GetSessionMessages :many
SELECT * FROM session_messages WHERE session_id = ? ORDER BY id ASC;
-- name: CountSessions :one
SELECT COUNT(*) FROM sessions;
+31
View File
@@ -0,0 +1,31 @@
-- name: RecordTokenUsage :one
INSERT INTO token_stats (session_id, turn, eval_count, prompt_tokens, model)
VALUES (?, ?, ?, ?, ?) RETURNING *;
-- name: GetSessionTokenStats :many
SELECT * FROM token_stats WHERE session_id = ? ORDER BY turn ASC;
-- name: GetSessionTotalTokens :one
SELECT
CAST(COALESCE(SUM(eval_count), 0) AS INTEGER) AS total_eval,
CAST(COALESCE(SUM(prompt_tokens), 0) AS INTEGER) AS total_prompt,
CAST(COUNT(*) AS INTEGER) AS turn_count
FROM token_stats WHERE session_id = ?;
-- name: RecordFileChange :one
INSERT INTO file_changes (session_id, file_path, tool_name, added, removed)
VALUES (?, ?, ?, ?, ?) RETURNING *;
-- name: GetSessionFileChanges :many
SELECT * FROM file_changes WHERE session_id = ? ORDER BY created_at ASC;
-- name: GetSessionFileChangeSummary :many
SELECT
file_path,
CAST(COALESCE(SUM(added), 0) AS INTEGER) AS total_added,
CAST(COALESCE(SUM(removed), 0) AS INTEGER) AS total_removed,
CAST(COUNT(*) AS INTEGER) AS change_count
FROM file_changes
WHERE session_id = ?
GROUP BY file_path
ORDER BY file_path ASC;
+206
View File
@@ -0,0 +1,206 @@
// Code generated by sqlc. DO NOT EDIT.
// versions:
// sqlc v1.30.0
// source: sessions.sql
package db
import (
"context"
)
const countSessions = `-- name: CountSessions :one
SELECT COUNT(*) FROM sessions
`
func (q *Queries) CountSessions(ctx context.Context) (int64, error) {
row := q.db.QueryRowContext(ctx, countSessions)
var count int64
err := row.Scan(&count)
return count, err
}
const createSession = `-- name: CreateSession :one
INSERT INTO sessions (title, model, mode) VALUES (?, ?, ?) RETURNING id, title, model, mode, created_at, updated_at
`
type CreateSessionParams struct {
Title string `json:"title"`
Model string `json:"model"`
Mode string `json:"mode"`
}
func (q *Queries) CreateSession(ctx context.Context, arg CreateSessionParams) (Session, error) {
row := q.db.QueryRowContext(ctx, createSession, arg.Title, arg.Model, arg.Mode)
var i Session
err := row.Scan(
&i.ID,
&i.Title,
&i.Model,
&i.Mode,
&i.CreatedAt,
&i.UpdatedAt,
)
return i, err
}
const createSessionMessage = `-- name: CreateSessionMessage :one
INSERT INTO session_messages (session_id, role, content, tool_name, tool_args, is_error, thinking)
VALUES (?, ?, ?, ?, ?, ?, ?) RETURNING id, session_id, role, content, tool_name, tool_args, is_error, thinking, created_at
`
type CreateSessionMessageParams struct {
SessionID int64 `json:"session_id"`
Role string `json:"role"`
Content string `json:"content"`
ToolName string `json:"tool_name"`
ToolArgs string `json:"tool_args"`
IsError int64 `json:"is_error"`
Thinking string `json:"thinking"`
}
func (q *Queries) CreateSessionMessage(ctx context.Context, arg CreateSessionMessageParams) (SessionMessage, error) {
row := q.db.QueryRowContext(ctx, createSessionMessage,
arg.SessionID,
arg.Role,
arg.Content,
arg.ToolName,
arg.ToolArgs,
arg.IsError,
arg.Thinking,
)
var i SessionMessage
err := row.Scan(
&i.ID,
&i.SessionID,
&i.Role,
&i.Content,
&i.ToolName,
&i.ToolArgs,
&i.IsError,
&i.Thinking,
&i.CreatedAt,
)
return i, err
}
const deleteSession = `-- name: DeleteSession :exec
DELETE FROM sessions WHERE id = ?
`
func (q *Queries) DeleteSession(ctx context.Context, id int64) error {
_, err := q.db.ExecContext(ctx, deleteSession, id)
return err
}
const getSession = `-- name: GetSession :one
SELECT id, title, model, mode, created_at, updated_at FROM sessions WHERE id = ?
`
func (q *Queries) GetSession(ctx context.Context, id int64) (Session, error) {
row := q.db.QueryRowContext(ctx, getSession, id)
var i Session
err := row.Scan(
&i.ID,
&i.Title,
&i.Model,
&i.Mode,
&i.CreatedAt,
&i.UpdatedAt,
)
return i, err
}
const getSessionMessages = `-- name: GetSessionMessages :many
SELECT id, session_id, role, content, tool_name, tool_args, is_error, thinking, created_at FROM session_messages WHERE session_id = ? ORDER BY id ASC
`
func (q *Queries) GetSessionMessages(ctx context.Context, sessionID int64) ([]SessionMessage, error) {
rows, err := q.db.QueryContext(ctx, getSessionMessages, sessionID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []SessionMessage{}
for rows.Next() {
var i SessionMessage
if err := rows.Scan(
&i.ID,
&i.SessionID,
&i.Role,
&i.Content,
&i.ToolName,
&i.ToolArgs,
&i.IsError,
&i.Thinking,
&i.CreatedAt,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const listSessions = `-- name: ListSessions :many
SELECT id, title, model, mode, created_at, updated_at FROM sessions ORDER BY updated_at DESC LIMIT ?
`
func (q *Queries) ListSessions(ctx context.Context, limit int64) ([]Session, error) {
rows, err := q.db.QueryContext(ctx, listSessions, limit)
if err != nil {
return nil, err
}
defer rows.Close()
items := []Session{}
for rows.Next() {
var i Session
if err := rows.Scan(
&i.ID,
&i.Title,
&i.Model,
&i.Mode,
&i.CreatedAt,
&i.UpdatedAt,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const updateSessionTimestamp = `-- name: UpdateSessionTimestamp :exec
UPDATE sessions SET updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now') WHERE id = ?
`
func (q *Queries) UpdateSessionTimestamp(ctx context.Context, id int64) error {
_, err := q.db.ExecContext(ctx, updateSessionTimestamp, id)
return err
}
const updateSessionTitle = `-- name: UpdateSessionTitle :exec
UPDATE sessions SET title = ?, updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now') WHERE id = ?
`
type UpdateSessionTitleParams struct {
Title string `json:"title"`
ID int64 `json:"id"`
}
func (q *Queries) UpdateSessionTitle(ctx context.Context, arg UpdateSessionTitleParams) error {
_, err := q.db.ExecContext(ctx, updateSessionTitle, arg.Title, arg.ID)
return err
}
+11
View File
@@ -0,0 +1,11 @@
version: "2"
sql:
- engine: "sqlite"
queries: "queries"
schema: "migrations"
gen:
go:
package: "db"
out: "."
emit_json_tags: true
emit_empty_slices: true
+216
View File
@@ -0,0 +1,216 @@
// Code generated by sqlc. DO NOT EDIT.
// versions:
// sqlc v1.30.0
// source: stats.sql
package db
import (
"context"
)
const getSessionFileChangeSummary = `-- name: GetSessionFileChangeSummary :many
SELECT
file_path,
CAST(COALESCE(SUM(added), 0) AS INTEGER) AS total_added,
CAST(COALESCE(SUM(removed), 0) AS INTEGER) AS total_removed,
CAST(COUNT(*) AS INTEGER) AS change_count
FROM file_changes
WHERE session_id = ?
GROUP BY file_path
ORDER BY file_path ASC
`
type GetSessionFileChangeSummaryRow struct {
FilePath string `json:"file_path"`
TotalAdded int64 `json:"total_added"`
TotalRemoved int64 `json:"total_removed"`
ChangeCount int64 `json:"change_count"`
}
func (q *Queries) GetSessionFileChangeSummary(ctx context.Context, sessionID int64) ([]GetSessionFileChangeSummaryRow, error) {
rows, err := q.db.QueryContext(ctx, getSessionFileChangeSummary, sessionID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []GetSessionFileChangeSummaryRow{}
for rows.Next() {
var i GetSessionFileChangeSummaryRow
if err := rows.Scan(
&i.FilePath,
&i.TotalAdded,
&i.TotalRemoved,
&i.ChangeCount,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getSessionFileChanges = `-- name: GetSessionFileChanges :many
SELECT id, session_id, file_path, tool_name, added, removed, created_at FROM file_changes WHERE session_id = ? ORDER BY created_at ASC
`
func (q *Queries) GetSessionFileChanges(ctx context.Context, sessionID int64) ([]FileChange, error) {
rows, err := q.db.QueryContext(ctx, getSessionFileChanges, sessionID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []FileChange{}
for rows.Next() {
var i FileChange
if err := rows.Scan(
&i.ID,
&i.SessionID,
&i.FilePath,
&i.ToolName,
&i.Added,
&i.Removed,
&i.CreatedAt,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getSessionTokenStats = `-- name: GetSessionTokenStats :many
SELECT id, session_id, turn, eval_count, prompt_tokens, model, created_at FROM token_stats WHERE session_id = ? ORDER BY turn ASC
`
func (q *Queries) GetSessionTokenStats(ctx context.Context, sessionID int64) ([]TokenStat, error) {
rows, err := q.db.QueryContext(ctx, getSessionTokenStats, sessionID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []TokenStat{}
for rows.Next() {
var i TokenStat
if err := rows.Scan(
&i.ID,
&i.SessionID,
&i.Turn,
&i.EvalCount,
&i.PromptTokens,
&i.Model,
&i.CreatedAt,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getSessionTotalTokens = `-- name: GetSessionTotalTokens :one
SELECT
CAST(COALESCE(SUM(eval_count), 0) AS INTEGER) AS total_eval,
CAST(COALESCE(SUM(prompt_tokens), 0) AS INTEGER) AS total_prompt,
CAST(COUNT(*) AS INTEGER) AS turn_count
FROM token_stats WHERE session_id = ?
`
type GetSessionTotalTokensRow struct {
TotalEval int64 `json:"total_eval"`
TotalPrompt int64 `json:"total_prompt"`
TurnCount int64 `json:"turn_count"`
}
func (q *Queries) GetSessionTotalTokens(ctx context.Context, sessionID int64) (GetSessionTotalTokensRow, error) {
row := q.db.QueryRowContext(ctx, getSessionTotalTokens, sessionID)
var i GetSessionTotalTokensRow
err := row.Scan(&i.TotalEval, &i.TotalPrompt, &i.TurnCount)
return i, err
}
const recordFileChange = `-- name: RecordFileChange :one
INSERT INTO file_changes (session_id, file_path, tool_name, added, removed)
VALUES (?, ?, ?, ?, ?) RETURNING id, session_id, file_path, tool_name, added, removed, created_at
`
type RecordFileChangeParams struct {
SessionID int64 `json:"session_id"`
FilePath string `json:"file_path"`
ToolName string `json:"tool_name"`
Added int64 `json:"added"`
Removed int64 `json:"removed"`
}
func (q *Queries) RecordFileChange(ctx context.Context, arg RecordFileChangeParams) (FileChange, error) {
row := q.db.QueryRowContext(ctx, recordFileChange,
arg.SessionID,
arg.FilePath,
arg.ToolName,
arg.Added,
arg.Removed,
)
var i FileChange
err := row.Scan(
&i.ID,
&i.SessionID,
&i.FilePath,
&i.ToolName,
&i.Added,
&i.Removed,
&i.CreatedAt,
)
return i, err
}
const recordTokenUsage = `-- name: RecordTokenUsage :one
INSERT INTO token_stats (session_id, turn, eval_count, prompt_tokens, model)
VALUES (?, ?, ?, ?, ?) RETURNING id, session_id, turn, eval_count, prompt_tokens, model, created_at
`
type RecordTokenUsageParams struct {
SessionID int64 `json:"session_id"`
Turn int64 `json:"turn"`
EvalCount int64 `json:"eval_count"`
PromptTokens int64 `json:"prompt_tokens"`
Model string `json:"model"`
}
func (q *Queries) RecordTokenUsage(ctx context.Context, arg RecordTokenUsageParams) (TokenStat, error) {
row := q.db.QueryRowContext(ctx, recordTokenUsage,
arg.SessionID,
arg.Turn,
arg.EvalCount,
arg.PromptTokens,
arg.Model,
)
var i TokenStat
err := row.Scan(
&i.ID,
&i.SessionID,
&i.Turn,
&i.EvalCount,
&i.PromptTokens,
&i.Model,
&i.CreatedAt,
)
return i, err
}
+71
View File
@@ -0,0 +1,71 @@
package db
import (
"database/sql"
"embed"
"fmt"
"os"
"path/filepath"
_ "modernc.org/sqlite"
)
//go:embed migrations/*.sql
var migrations embed.FS
type Store struct {
*Queries
db *sql.DB
}
func Open() (*Store, error) {
home, err := os.UserHomeDir()
if err != nil {
return nil, fmt.Errorf("home dir: %w", err)
}
dir := filepath.Join(home, ".config", "ai-agent")
if err := os.MkdirAll(dir, 0o755); err != nil {
return nil, fmt.Errorf("create config dir: %w", err)
}
return OpenPath(filepath.Join(dir, "ai-agent.db"))
}
func OpenPath(path string) (*Store, error) {
conn, err := sql.Open("sqlite", path+"?_journal_mode=WAL&_busy_timeout=5000&_foreign_keys=ON")
if err != nil {
return nil, fmt.Errorf("open db: %w", err)
}
if err := runMigrations(conn); err != nil {
conn.Close()
return nil, fmt.Errorf("migrations: %w", err)
}
return &Store{Queries: New(conn), db: conn}, nil
}
func (s *Store) Close() error {
return s.db.Close()
}
func (s *Store) DB() *sql.DB {
return s.db
}
func runMigrations(conn *sql.DB) error {
entries, err := migrations.ReadDir("migrations")
if err != nil {
return fmt.Errorf("read migrations dir: %w", err)
}
for _, entry := range entries {
if entry.IsDir() {
continue
}
data, err := migrations.ReadFile("migrations/" + entry.Name())
if err != nil {
return fmt.Errorf("read migration %s: %w", entry.Name(), err)
}
if _, err := conn.Exec(string(data)); err != nil {
return fmt.Errorf("exec migration %s: %w", entry.Name(), err)
}
}
return nil
}
+271
View File
@@ -0,0 +1,271 @@
package db
import (
"context"
"os"
"path/filepath"
"testing"
)
func testStore(t *testing.T) *Store {
t.Helper()
dir := t.TempDir()
s, err := OpenPath(filepath.Join(dir, "test.db"))
if err != nil {
t.Fatalf("open store: %v", err)
}
t.Cleanup(func() { s.Close() })
return s
}
func TestOpenAndMigrate(t *testing.T) {
s := testStore(t)
// Verify the store is functional by counting sessions.
ctx := context.Background()
count, err := s.CountSessions(ctx)
if err != nil {
t.Fatalf("count sessions: %v", err)
}
if count != 0 {
t.Fatalf("expected 0 sessions, got %d", count)
}
}
func TestSessionCRUD(t *testing.T) {
s := testStore(t)
ctx := context.Background()
// Create.
sess, err := s.CreateSession(ctx, CreateSessionParams{
Title: "Test Session",
Model: "qwen3.5:4b",
Mode: "BUILD",
})
if err != nil {
t.Fatalf("create session: %v", err)
}
if sess.Title != "Test Session" {
t.Fatalf("expected title 'Test Session', got %q", sess.Title)
}
// Read.
got, err := s.GetSession(ctx, sess.ID)
if err != nil {
t.Fatalf("get session: %v", err)
}
if got.Model != "qwen3.5:4b" {
t.Fatalf("expected model 'qwen3.5:4b', got %q", got.Model)
}
// Create message.
msg, err := s.CreateSessionMessage(ctx, CreateSessionMessageParams{
SessionID: sess.ID,
Role: "user",
Content: "Hello world",
})
if err != nil {
t.Fatalf("create message: %v", err)
}
if msg.Content != "Hello world" {
t.Fatalf("expected content 'Hello world', got %q", msg.Content)
}
// List messages.
msgs, err := s.GetSessionMessages(ctx, sess.ID)
if err != nil {
t.Fatalf("get messages: %v", err)
}
if len(msgs) != 1 {
t.Fatalf("expected 1 message, got %d", len(msgs))
}
// Delete.
if err := s.DeleteSession(ctx, sess.ID); err != nil {
t.Fatalf("delete session: %v", err)
}
count, err := s.CountSessions(ctx)
if err != nil {
t.Fatalf("count: %v", err)
}
if count != 0 {
t.Fatalf("expected 0 after delete, got %d", count)
}
}
func TestToolPermissions(t *testing.T) {
s := testStore(t)
ctx := context.Background()
// Upsert.
perm, err := s.UpsertToolPermission(ctx, UpsertToolPermissionParams{
ToolName: "bash",
Policy: "allow",
})
if err != nil {
t.Fatalf("upsert permission: %v", err)
}
if perm.Policy != "allow" {
t.Fatalf("expected policy 'allow', got %q", perm.Policy)
}
// Update via upsert.
perm2, err := s.UpsertToolPermission(ctx, UpsertToolPermissionParams{
ToolName: "bash",
Policy: "deny",
})
if err != nil {
t.Fatalf("upsert update: %v", err)
}
if perm2.Policy != "deny" {
t.Fatalf("expected policy 'deny', got %q", perm2.Policy)
}
// List.
perms, err := s.ListToolPermissions(ctx)
if err != nil {
t.Fatalf("list permissions: %v", err)
}
if len(perms) != 1 {
t.Fatalf("expected 1 permission, got %d", len(perms))
}
// Reset.
if err := s.ResetToolPermissions(ctx); err != nil {
t.Fatalf("reset: %v", err)
}
perms, err = s.ListToolPermissions(ctx)
if err != nil {
t.Fatalf("list after reset: %v", err)
}
if len(perms) != 0 {
t.Fatalf("expected 0 after reset, got %d", len(perms))
}
}
func TestTokenStats(t *testing.T) {
s := testStore(t)
ctx := context.Background()
sess, err := s.CreateSession(ctx, CreateSessionParams{
Title: "Stats Test",
Model: "qwen3.5:4b",
Mode: "ASK",
})
if err != nil {
t.Fatalf("create session: %v", err)
}
// Record usage.
_, err = s.RecordTokenUsage(ctx, RecordTokenUsageParams{
SessionID: sess.ID,
Turn: 1,
EvalCount: 100,
PromptTokens: 500,
Model: "qwen3.5:4b",
})
if err != nil {
t.Fatalf("record usage: %v", err)
}
_, err = s.RecordTokenUsage(ctx, RecordTokenUsageParams{
SessionID: sess.ID,
Turn: 2,
EvalCount: 200,
PromptTokens: 600,
Model: "qwen3.5:4b",
})
if err != nil {
t.Fatalf("record usage 2: %v", err)
}
// Get totals.
totals, err := s.GetSessionTotalTokens(ctx, sess.ID)
if err != nil {
t.Fatalf("get totals: %v", err)
}
if totals.TotalEval != 300 {
t.Fatalf("expected total_eval 300, got %v", totals.TotalEval)
}
if totals.TotalPrompt != 1100 {
t.Fatalf("expected total_prompt 1100, got %v", totals.TotalPrompt)
}
if totals.TurnCount != 2 {
t.Fatalf("expected turn_count 2, got %v", totals.TurnCount)
}
}
func TestFileChanges(t *testing.T) {
s := testStore(t)
ctx := context.Background()
sess, err := s.CreateSession(ctx, CreateSessionParams{
Title: "Changes Test",
Model: "qwen3.5:4b",
Mode: "BUILD",
})
if err != nil {
t.Fatalf("create session: %v", err)
}
_, err = s.RecordFileChange(ctx, RecordFileChangeParams{
SessionID: sess.ID,
FilePath: "main.go",
ToolName: "write_file",
Added: 10,
Removed: 3,
})
if err != nil {
t.Fatalf("record change: %v", err)
}
_, err = s.RecordFileChange(ctx, RecordFileChangeParams{
SessionID: sess.ID,
FilePath: "main.go",
ToolName: "write_file",
Added: 5,
Removed: 2,
})
if err != nil {
t.Fatalf("record change 2: %v", err)
}
summary, err := s.GetSessionFileChangeSummary(ctx, sess.ID)
if err != nil {
t.Fatalf("get summary: %v", err)
}
if len(summary) != 1 {
t.Fatalf("expected 1 file, got %d", len(summary))
}
if summary[0].TotalAdded != 15 {
t.Fatalf("expected 15 added, got %d", summary[0].TotalAdded)
}
}
func TestDoubleOpen(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "test.db")
s1, err := OpenPath(path)
if err != nil {
t.Fatalf("first open: %v", err)
}
defer s1.Close()
// Running migrations again should be idempotent (IF NOT EXISTS).
s2, err := OpenPath(path)
if err != nil {
t.Fatalf("second open: %v", err)
}
defer s2.Close()
}
func TestOpenDefault(t *testing.T) {
if os.Getenv("CI") != "" {
t.Skip("skip in CI to avoid side effects")
}
s, err := Open()
if err != nil {
t.Fatalf("open default: %v", err)
}
s.Close()
}