demo
This commit is contained in:
601
db/db.go
Normal file
601
db/db.go
Normal file
@@ -0,0 +1,601 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
const dbPath = "/Users/rdarius/.nyxtex/go-coder.db"
|
||||
|
||||
type Group struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Location string `json:"location"`
|
||||
TotalTokens int `json:"total_tokens,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type Conversation struct {
|
||||
ID string `json:"id"`
|
||||
Title string `json:"title"`
|
||||
GroupID *string `json:"group_id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
MessageCount int `json:"message_count"`
|
||||
TotalTokens int `json:"total_tokens,omitempty"`
|
||||
ContextSummary string `json:"-"`
|
||||
ContextSummaryOrder int `json:"-"`
|
||||
}
|
||||
|
||||
type Message struct {
|
||||
ID int64 `json:"id"`
|
||||
ConvID string `json:"conversation_id"`
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
Kind string `json:"kind"`
|
||||
Name string `json:"name"`
|
||||
Agent string `json:"agent"`
|
||||
TotalTokens int `json:"total_tokens,omitempty"`
|
||||
Order int `json:"order"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type Store struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewStore() (*Store, error) {
|
||||
db, err := sql.Open("sqlite", dbPath+"?_busy_timeout=5000&_journal=WAL")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open database: %w", err)
|
||||
}
|
||||
|
||||
if err := db.Ping(); err != nil {
|
||||
return nil, fmt.Errorf("failed to ping database: %w", err)
|
||||
}
|
||||
|
||||
store := &Store{db: db}
|
||||
if err := store.createTables(); err != nil {
|
||||
return nil, fmt.Errorf("failed to create tables: %w", err)
|
||||
}
|
||||
|
||||
return store, nil
|
||||
}
|
||||
|
||||
func generateID() string {
|
||||
return uuid.New().String()
|
||||
}
|
||||
|
||||
func (s *Store) createTables() error {
|
||||
_, err := s.db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS groups (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS conversations (
|
||||
id TEXT PRIMARY KEY,
|
||||
title TEXT NOT NULL,
|
||||
group_id TEXT,
|
||||
context_summary TEXT NOT NULL DEFAULT '',
|
||||
context_summary_order INTEGER NOT NULL DEFAULT 0,
|
||||
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
FOREIGN KEY (group_id) REFERENCES groups(id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
conversation_id TEXT NOT NULL,
|
||||
role TEXT NOT NULL CHECK(role IN ('user', 'assistant', 'system')),
|
||||
content TEXT NOT NULL,
|
||||
total_tokens INTEGER NOT NULL DEFAULT 0,
|
||||
"order" INTEGER NOT NULL,
|
||||
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
applied_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_messages_conv_order ON messages(conversation_id, "order");
|
||||
CREATE INDEX IF NOT EXISTS idx_messages_conv_id ON messages(conversation_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_conversations_group_id ON conversations(group_id);
|
||||
`)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.applyMigrations()
|
||||
}
|
||||
|
||||
func (s *Store) applyMigrations() error {
|
||||
migrations := []struct {
|
||||
version int
|
||||
sql string
|
||||
}{
|
||||
{
|
||||
version: 1,
|
||||
sql: `ALTER TABLE groups ADD COLUMN location TEXT NOT NULL DEFAULT ''`,
|
||||
},
|
||||
{
|
||||
version: 2,
|
||||
sql: `ALTER TABLE messages ADD COLUMN kind TEXT NOT NULL DEFAULT 'message'`,
|
||||
},
|
||||
{
|
||||
version: 3,
|
||||
sql: `ALTER TABLE messages ADD COLUMN name TEXT NOT NULL DEFAULT ''`,
|
||||
},
|
||||
{
|
||||
version: 4,
|
||||
sql: `ALTER TABLE messages ADD COLUMN agent TEXT NOT NULL DEFAULT ''`,
|
||||
},
|
||||
{
|
||||
version: 5,
|
||||
sql: `ALTER TABLE messages ADD COLUMN total_tokens INTEGER NOT NULL DEFAULT 0`,
|
||||
},
|
||||
{
|
||||
version: 6,
|
||||
sql: `ALTER TABLE conversations ADD COLUMN context_summary TEXT NOT NULL DEFAULT ''`,
|
||||
},
|
||||
{
|
||||
version: 7,
|
||||
sql: `ALTER TABLE conversations ADD COLUMN context_summary_order INTEGER NOT NULL DEFAULT 0`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, m := range migrations {
|
||||
applied, err := s.isMigrationApplied(m.version)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !applied {
|
||||
if m.version == 5 {
|
||||
exists, err := s.columnExists("messages", "total_tokens")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
_, err = s.db.Exec(
|
||||
"INSERT INTO schema_migrations (version, applied_at) VALUES (?, ?)",
|
||||
m.version, time.Now().UTC(),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to record migration %d: %w", m.version, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
if m.version == 6 {
|
||||
exists, err := s.columnExists("conversations", "context_summary")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
_, err = s.db.Exec(
|
||||
"INSERT INTO schema_migrations (version, applied_at) VALUES (?, ?)",
|
||||
m.version, time.Now().UTC(),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to record migration %d: %w", m.version, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
if m.version == 7 {
|
||||
exists, err := s.columnExists("conversations", "context_summary_order")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
_, err = s.db.Exec(
|
||||
"INSERT INTO schema_migrations (version, applied_at) VALUES (?, ?)",
|
||||
m.version, time.Now().UTC(),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to record migration %d: %w", m.version, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
_, err = s.db.Exec(m.sql)
|
||||
if err != nil {
|
||||
return fmt.Errorf("migration %d failed: %w", m.version, err)
|
||||
}
|
||||
_, err = s.db.Exec(
|
||||
"INSERT INTO schema_migrations (version, applied_at) VALUES (?, ?)",
|
||||
m.version, time.Now().UTC(),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to record migration %d: %w", m.version, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) isMigrationApplied(version int) (bool, error) {
|
||||
var count int
|
||||
err := s.db.QueryRow(
|
||||
"SELECT COUNT(*) FROM schema_migrations WHERE version = ?", version,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s *Store) columnExists(tableName, columnName string) (bool, error) {
|
||||
rows, err := s.db.Query(fmt.Sprintf("PRAGMA table_info(%s)", tableName))
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var cid int
|
||||
var name, colType string
|
||||
var notNull int
|
||||
var dfltValue sql.NullString
|
||||
var pk int
|
||||
for rows.Next() {
|
||||
if err := rows.Scan(&cid, &name, &colType, ¬Null, &dfltValue, &pk); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if name == columnName {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (s *Store) CreateGroup(name string) (*Group, error) {
|
||||
id := generateID()
|
||||
now := time.Now().UTC()
|
||||
_, err := s.db.Exec(
|
||||
"INSERT INTO groups (id, name, created_at, updated_at) VALUES (?, ?, ?, ?)",
|
||||
id, name, now, now,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create group: %w", err)
|
||||
}
|
||||
return &Group{
|
||||
ID: id,
|
||||
Name: name,
|
||||
Location: "",
|
||||
TotalTokens: 0,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Store) GetGroup(id string) (*Group, error) {
|
||||
var group Group
|
||||
err := s.db.QueryRow(
|
||||
`SELECT g.id, g.name, g.location, g.created_at, g.updated_at,
|
||||
COALESCE((
|
||||
SELECT SUM(m.total_tokens)
|
||||
FROM conversations c
|
||||
JOIN messages m ON c.id = m.conversation_id
|
||||
WHERE c.group_id = g.id
|
||||
), 0) AS total_tokens
|
||||
FROM groups g
|
||||
WHERE g.id = ?`, id,
|
||||
).Scan(&group.ID, &group.Name, &group.Location, &group.CreatedAt, &group.UpdatedAt, &group.TotalTokens)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get group: %w", err)
|
||||
}
|
||||
return &group, nil
|
||||
}
|
||||
|
||||
func (s *Store) ListGroups() ([]Group, error) {
|
||||
rows, err := s.db.Query(
|
||||
`SELECT g.id, g.name, g.location, g.created_at, g.updated_at,
|
||||
COALESCE((
|
||||
SELECT SUM(m.total_tokens)
|
||||
FROM conversations c
|
||||
JOIN messages m ON c.id = m.conversation_id
|
||||
WHERE c.group_id = g.id
|
||||
), 0) AS total_tokens
|
||||
FROM groups g
|
||||
ORDER BY g.updated_at DESC`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list groups: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var groups []Group
|
||||
for rows.Next() {
|
||||
var g Group
|
||||
if err := rows.Scan(&g.ID, &g.Name, &g.Location, &g.CreatedAt, &g.UpdatedAt, &g.TotalTokens); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan group: %w", err)
|
||||
}
|
||||
groups = append(groups, g)
|
||||
}
|
||||
return groups, nil
|
||||
}
|
||||
|
||||
func (s *Store) UpdateGroupTitle(id, name string) error {
|
||||
_, err := s.db.Exec(
|
||||
"UPDATE groups SET name = ?, updated_at = ? WHERE id = ?",
|
||||
name, time.Now().UTC(), id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) DeleteGroup(id string) error {
|
||||
_, err := s.db.Exec("DELETE FROM groups WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) SetGroupLocation(id, location string) error {
|
||||
_, err := s.db.Exec(
|
||||
"UPDATE groups SET location = ?, updated_at = ? WHERE id = ?",
|
||||
location, time.Now().UTC(), id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) UpdateConversationGroup(id, groupID string) error {
|
||||
_, err := s.db.Exec(
|
||||
"UPDATE conversations SET group_id = ?, updated_at = ? WHERE id = ? AND (group_id IS NULL OR group_id = '')",
|
||||
groupID, time.Now().UTC(), id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) HasGroupLocation(id string) (bool, error) {
|
||||
var count int
|
||||
err := s.db.QueryRow(
|
||||
"SELECT COUNT(*) FROM groups WHERE id = ? AND location != ''", id,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s *Store) CreateConversation(title string, groupID *string) (*Conversation, error) {
|
||||
id := generateID()
|
||||
now := time.Now().UTC()
|
||||
var groupIDVal interface{} = nil
|
||||
if groupID != nil {
|
||||
groupIDVal = *groupID
|
||||
}
|
||||
_, err := s.db.Exec(
|
||||
"INSERT INTO conversations (id, title, group_id, created_at, updated_at) VALUES (?, ?, ?, ?, ?)",
|
||||
id, title, groupIDVal, now, now,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create conversation: %w", err)
|
||||
}
|
||||
return &Conversation{
|
||||
ID: id,
|
||||
Title: title,
|
||||
GroupID: groupID,
|
||||
TotalTokens: 0,
|
||||
ContextSummary: "",
|
||||
ContextSummaryOrder: 0,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Store) GetConversation(id string) (*Conversation, error) {
|
||||
var conv Conversation
|
||||
var groupID sql.NullString
|
||||
err := s.db.QueryRow(
|
||||
`SELECT c.id, c.title, c.group_id, c.created_at, c.updated_at,
|
||||
COALESCE((SELECT COUNT(*) FROM messages m WHERE m.conversation_id = c.id), 0) AS message_count,
|
||||
COALESCE((SELECT SUM(m.total_tokens) FROM messages m WHERE m.conversation_id = c.id), 0) AS total_tokens,
|
||||
COALESCE(c.context_summary, '') AS context_summary,
|
||||
COALESCE(c.context_summary_order, 0) AS context_summary_order
|
||||
FROM conversations c
|
||||
WHERE c.id = ?`, id,
|
||||
).Scan(&conv.ID, &conv.Title, &groupID, &conv.CreatedAt, &conv.UpdatedAt, &conv.MessageCount, &conv.TotalTokens, &conv.ContextSummary, &conv.ContextSummaryOrder)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get conversation: %w", err)
|
||||
}
|
||||
if groupID.Valid {
|
||||
conv.GroupID = &groupID.String
|
||||
}
|
||||
return &conv, nil
|
||||
}
|
||||
|
||||
func (s *Store) ListConversations() ([]Conversation, error) {
|
||||
rows, err := s.db.Query(`
|
||||
SELECT c.id, c.title, c.group_id, c.created_at, c.updated_at,
|
||||
COALESCE((SELECT COUNT(*) FROM messages m WHERE m.conversation_id = c.id), 0) AS message_count,
|
||||
COALESCE((SELECT SUM(m.total_tokens) FROM messages m WHERE m.conversation_id = c.id), 0) AS total_tokens,
|
||||
COALESCE(c.context_summary, '') AS context_summary,
|
||||
COALESCE(c.context_summary_order, 0) AS context_summary_order
|
||||
FROM conversations c
|
||||
ORDER BY c.updated_at DESC
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list conversations: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var conversations []Conversation
|
||||
for rows.Next() {
|
||||
var conv Conversation
|
||||
var groupID sql.NullString
|
||||
if err := rows.Scan(&conv.ID, &conv.Title, &groupID, &conv.CreatedAt, &conv.UpdatedAt, &conv.MessageCount, &conv.TotalTokens, &conv.ContextSummary, &conv.ContextSummaryOrder); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan conversation: %w", err)
|
||||
}
|
||||
if groupID.Valid {
|
||||
conv.GroupID = &groupID.String
|
||||
}
|
||||
conversations = append(conversations, conv)
|
||||
}
|
||||
return conversations, nil
|
||||
}
|
||||
|
||||
type GroupedConversations struct {
|
||||
Groups []GroupWithConversations `json:"groups"`
|
||||
Ungrouped []Conversation `json:"ungrouped"`
|
||||
}
|
||||
|
||||
type GroupWithConversations struct {
|
||||
Group
|
||||
Conversations []Conversation `json:"conversations"`
|
||||
}
|
||||
|
||||
func (s *Store) ListGroupedConversations() (*GroupedConversations, error) {
|
||||
groups, err := s.ListGroups()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
grouped := &GroupedConversations{
|
||||
Groups: make([]GroupWithConversations, 0),
|
||||
Ungrouped: make([]Conversation, 0),
|
||||
}
|
||||
|
||||
for _, g := range groups {
|
||||
rows, err := s.db.Query(`
|
||||
SELECT c.id, c.title, c.group_id, c.created_at, c.updated_at,
|
||||
COALESCE((SELECT COUNT(*) FROM messages m WHERE m.conversation_id = c.id), 0) AS message_count,
|
||||
COALESCE((SELECT SUM(m.total_tokens) FROM messages m WHERE m.conversation_id = c.id), 0) AS total_tokens,
|
||||
COALESCE(c.context_summary, '') AS context_summary,
|
||||
COALESCE(c.context_summary_order, 0) AS context_summary_order
|
||||
FROM conversations c
|
||||
WHERE c.group_id = ?
|
||||
ORDER BY c.updated_at DESC
|
||||
`, g.ID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list conversations for group %s: %w", g.ID, err)
|
||||
}
|
||||
|
||||
var convs []Conversation
|
||||
for rows.Next() {
|
||||
var conv Conversation
|
||||
var groupID sql.NullString
|
||||
if err := rows.Scan(&conv.ID, &conv.Title, &groupID, &conv.CreatedAt, &conv.UpdatedAt, &conv.MessageCount, &conv.TotalTokens, &conv.ContextSummary, &conv.ContextSummaryOrder); err != nil {
|
||||
rows.Close()
|
||||
return nil, fmt.Errorf("failed to scan conversation: %w", err)
|
||||
}
|
||||
convs = append(convs, conv)
|
||||
}
|
||||
rows.Close()
|
||||
|
||||
grouped.Groups = append(grouped.Groups, GroupWithConversations{
|
||||
Group: g,
|
||||
Conversations: convs,
|
||||
})
|
||||
}
|
||||
|
||||
ungroupedRows, err := s.db.Query(`
|
||||
SELECT c.id, c.title, c.group_id, c.created_at, c.updated_at,
|
||||
COALESCE((SELECT COUNT(*) FROM messages m WHERE m.conversation_id = c.id), 0) AS message_count,
|
||||
COALESCE((SELECT SUM(m.total_tokens) FROM messages m WHERE m.conversation_id = c.id), 0) AS total_tokens,
|
||||
COALESCE(c.context_summary, '') AS context_summary,
|
||||
COALESCE(c.context_summary_order, 0) AS context_summary_order
|
||||
FROM conversations c
|
||||
WHERE c.group_id IS NULL
|
||||
ORDER BY c.updated_at DESC
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list ungrouped conversations: %w", err)
|
||||
}
|
||||
defer ungroupedRows.Close()
|
||||
|
||||
for ungroupedRows.Next() {
|
||||
var conv Conversation
|
||||
var groupID sql.NullString
|
||||
if err := ungroupedRows.Scan(&conv.ID, &conv.Title, &groupID, &conv.CreatedAt, &conv.UpdatedAt, &conv.MessageCount, &conv.TotalTokens, &conv.ContextSummary, &conv.ContextSummaryOrder); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan ungrouped conversation: %w", err)
|
||||
}
|
||||
grouped.Ungrouped = append(grouped.Ungrouped, conv)
|
||||
}
|
||||
|
||||
return grouped, nil
|
||||
}
|
||||
|
||||
func (s *Store) UpdateConversationTitle(id, title string) error {
|
||||
_, err := s.db.Exec(
|
||||
"UPDATE conversations SET title = ?, updated_at = ? WHERE id = ?",
|
||||
title, time.Now().UTC(), id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) UpdateConversationContextSummary(id, summary string, summaryOrder int) error {
|
||||
_, err := s.db.Exec(
|
||||
"UPDATE conversations SET context_summary = ?, context_summary_order = ?, updated_at = ? WHERE id = ?",
|
||||
summary, summaryOrder, time.Now().UTC(), id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) DeleteConversation(id string) error {
|
||||
_, err := s.db.Exec("DELETE FROM conversations WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) AddMessage(convID, role, content string, order int, totalTokens int) error {
|
||||
return s.AddMessageWithMeta(convID, role, content, order, "message", "", "", totalTokens)
|
||||
}
|
||||
|
||||
func (s *Store) AddMessageWithMeta(convID, role, content string, order int, kind, name, agent string, totalTokens int) error {
|
||||
if kind == "" {
|
||||
kind = "message"
|
||||
}
|
||||
_, err := s.db.Exec(
|
||||
"INSERT INTO messages (conversation_id, role, content, total_tokens, kind, name, agent, \"order\", created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
convID, role, content, totalTokens, kind, name, agent, order, time.Now().UTC(),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to add message: %w", err)
|
||||
}
|
||||
_, err = s.db.Exec(
|
||||
"UPDATE conversations SET updated_at = ? WHERE id = ?",
|
||||
time.Now().UTC(), convID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) GetMessages(convID string) ([]Message, error) {
|
||||
rows, err := s.db.Query(
|
||||
"SELECT id, conversation_id, role, content, COALESCE(kind, 'message'), COALESCE(name, ''), COALESCE(agent, ''), COALESCE(total_tokens, 0), \"order\", created_at FROM messages WHERE conversation_id = ? ORDER BY \"order\" ASC",
|
||||
convID,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get messages: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
messages := make([]Message, 0)
|
||||
for rows.Next() {
|
||||
var msg Message
|
||||
if err := rows.Scan(&msg.ID, &msg.ConvID, &msg.Role, &msg.Content, &msg.Kind, &msg.Name, &msg.Agent, &msg.TotalTokens, &msg.Order, &msg.CreatedAt); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan message: %w", err)
|
||||
}
|
||||
messages = append(messages, msg)
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
func (s *Store) DeleteMessages(convID string) error {
|
||||
_, err := s.db.Exec("DELETE FROM messages WHERE conversation_id = ?", convID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) Close() error {
|
||||
return s.db.Close()
|
||||
}
|
||||
211
db/db_test.go
Normal file
211
db/db_test.go
Normal file
@@ -0,0 +1,211 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func newTempStore(t *testing.T) *Store {
|
||||
t.Helper()
|
||||
|
||||
dbFile := filepath.Join(t.TempDir(), "test.db")
|
||||
conn, err := sql.Open("sqlite", dbFile+"?_busy_timeout=5000&_journal=WAL")
|
||||
if err != nil {
|
||||
t.Fatalf("sql.Open returned error: %v", err)
|
||||
}
|
||||
|
||||
store := &Store{db: conn}
|
||||
if err := store.createTables(); err != nil {
|
||||
t.Fatalf("createTables returned error: %v", err)
|
||||
}
|
||||
|
||||
t.Cleanup(func() {
|
||||
_ = conn.Close()
|
||||
})
|
||||
|
||||
return store
|
||||
}
|
||||
|
||||
func TestSetGroupLocationUpdatesExistingLocation(t *testing.T) {
|
||||
store := newTempStore(t)
|
||||
|
||||
group, err := store.CreateGroup("Workspace")
|
||||
if err != nil {
|
||||
t.Fatalf("CreateGroup returned error: %v", err)
|
||||
}
|
||||
|
||||
if err := store.SetGroupLocation(group.ID, "/tmp/one"); err != nil {
|
||||
t.Fatalf("SetGroupLocation returned error: %v", err)
|
||||
}
|
||||
|
||||
updated, err := store.GetGroup(group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetGroup returned error: %v", err)
|
||||
}
|
||||
if updated.Location != "/tmp/one" {
|
||||
t.Fatalf("location after first update = %q, want %q", updated.Location, "/tmp/one")
|
||||
}
|
||||
|
||||
if err := store.SetGroupLocation(group.ID, "/tmp/two"); err != nil {
|
||||
t.Fatalf("SetGroupLocation returned error: %v", err)
|
||||
}
|
||||
|
||||
updated, err = store.GetGroup(group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetGroup returned error: %v", err)
|
||||
}
|
||||
if updated.Location != "/tmp/two" {
|
||||
t.Fatalf("location after second update = %q, want %q", updated.Location, "/tmp/two")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateConversationGroupAttachesUngroupedConversation(t *testing.T) {
|
||||
store := newTempStore(t)
|
||||
|
||||
group, err := store.CreateGroup("Workspace")
|
||||
if err != nil {
|
||||
t.Fatalf("CreateGroup returned error: %v", err)
|
||||
}
|
||||
|
||||
conv, err := store.CreateConversation("Chat", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateConversation returned error: %v", err)
|
||||
}
|
||||
|
||||
if err := store.UpdateConversationGroup(conv.ID, group.ID); err != nil {
|
||||
t.Fatalf("UpdateConversationGroup returned error: %v", err)
|
||||
}
|
||||
|
||||
updated, err := store.GetConversation(conv.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetConversation returned error: %v", err)
|
||||
}
|
||||
if updated.GroupID == nil {
|
||||
t.Fatal("conversation group ID is nil, want a group ID")
|
||||
}
|
||||
if got, want := *updated.GroupID, group.ID; got != want {
|
||||
t.Fatalf("conversation group ID = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddMessageWithMetaPersistsKindNameAndAgent(t *testing.T) {
|
||||
store := newTempStore(t)
|
||||
|
||||
conv, err := store.CreateConversation("Chat", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateConversation returned error: %v", err)
|
||||
}
|
||||
|
||||
if err := store.AddMessageWithMeta(conv.ID, "assistant", "Thinking", 1, "thinking", "", "Planner", 42); err != nil {
|
||||
t.Fatalf("AddMessageWithMeta returned error: %v", err)
|
||||
}
|
||||
|
||||
messages, err := store.GetMessages(conv.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetMessages returned error: %v", err)
|
||||
}
|
||||
if len(messages) != 1 {
|
||||
t.Fatalf("message count = %d, want 1", len(messages))
|
||||
}
|
||||
if got, want := messages[0].Kind, "thinking"; got != want {
|
||||
t.Fatalf("message kind = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := messages[0].Name, ""; got != want {
|
||||
t.Fatalf("message name = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := messages[0].Agent, "Planner"; got != want {
|
||||
t.Fatalf("message agent = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := messages[0].TotalTokens, 42; got != want {
|
||||
t.Fatalf("message total tokens = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConversationAndGroupTokenTotalsAggregate(t *testing.T) {
|
||||
store := newTempStore(t)
|
||||
|
||||
group, err := store.CreateGroup("Workspace")
|
||||
if err != nil {
|
||||
t.Fatalf("CreateGroup returned error: %v", err)
|
||||
}
|
||||
|
||||
conv1, err := store.CreateConversation("One", &group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateConversation returned error: %v", err)
|
||||
}
|
||||
conv2, err := store.CreateConversation("Two", &group.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateConversation returned error: %v", err)
|
||||
}
|
||||
ungrouped, err := store.CreateConversation("Three", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateConversation returned error: %v", err)
|
||||
}
|
||||
|
||||
if err := store.AddMessageWithMeta(conv1.ID, "assistant", "A", 1, "message", "", "Planner", 11); err != nil {
|
||||
t.Fatalf("AddMessageWithMeta returned error: %v", err)
|
||||
}
|
||||
if err := store.AddMessageWithMeta(conv1.ID, "assistant", "B", 2, "message", "", "Planner", 7); err != nil {
|
||||
t.Fatalf("AddMessageWithMeta returned error: %v", err)
|
||||
}
|
||||
if err := store.AddMessageWithMeta(conv2.ID, "assistant", "C", 1, "message", "", "Planner", 5); err != nil {
|
||||
t.Fatalf("AddMessageWithMeta returned error: %v", err)
|
||||
}
|
||||
if err := store.AddMessageWithMeta(ungrouped.ID, "assistant", "D", 1, "message", "", "Planner", 3); err != nil {
|
||||
t.Fatalf("AddMessageWithMeta returned error: %v", err)
|
||||
}
|
||||
|
||||
conversations, err := store.ListConversations()
|
||||
if err != nil {
|
||||
t.Fatalf("ListConversations returned error: %v", err)
|
||||
}
|
||||
totalsByID := make(map[string]int, len(conversations))
|
||||
for _, conv := range conversations {
|
||||
totalsByID[conv.ID] = conv.TotalTokens
|
||||
}
|
||||
if got, want := totalsByID[conv1.ID], 18; got != want {
|
||||
t.Fatalf("conversation one total tokens = %d, want %d", got, want)
|
||||
}
|
||||
if got, want := totalsByID[conv2.ID], 5; got != want {
|
||||
t.Fatalf("conversation two total tokens = %d, want %d", got, want)
|
||||
}
|
||||
if got, want := totalsByID[ungrouped.ID], 3; got != want {
|
||||
t.Fatalf("ungrouped conversation total tokens = %d, want %d", got, want)
|
||||
}
|
||||
|
||||
groups, err := store.ListGroups()
|
||||
if err != nil {
|
||||
t.Fatalf("ListGroups returned error: %v", err)
|
||||
}
|
||||
if len(groups) != 1 {
|
||||
t.Fatalf("group count = %d, want 1", len(groups))
|
||||
}
|
||||
if got, want := groups[0].TotalTokens, 23; got != want {
|
||||
t.Fatalf("group total tokens = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateConversationContextSummaryPersistsSummaryAndOrder(t *testing.T) {
|
||||
store := newTempStore(t)
|
||||
|
||||
conv, err := store.CreateConversation("Chat", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateConversation returned error: %v", err)
|
||||
}
|
||||
|
||||
if err := store.UpdateConversationContextSummary(conv.ID, "compressed memory", 12); err != nil {
|
||||
t.Fatalf("UpdateConversationContextSummary returned error: %v", err)
|
||||
}
|
||||
|
||||
updated, err := store.GetConversation(conv.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetConversation returned error: %v", err)
|
||||
}
|
||||
if got, want := updated.ContextSummary, "compressed memory"; got != want {
|
||||
t.Fatalf("context summary = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := updated.ContextSummaryOrder, 12; got != want {
|
||||
t.Fatalf("context summary order = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user