fix(seahorse): wrap DeleteMessagesAfterID and appendContextItems in transactions

- DeleteMessagesAfterID: wrap all DELETE operations in a transaction for
  atomicity, remove redundant manual FTS delete (handled by trigger)
- appendContextItems: use transaction to fix read-then-write race condition
- Add GetMaxOrdinalTx and resolveItemTokenCountTx for transaction-scoped queries
- Remove unused resolveItemTokenCount function

Fixes PR review issues 6 and 7.
This commit is contained in:
Liu Yuan 2026-04-04 12:37:30 +08:00
parent 9ace45673a
commit 02c040ba9e

View file

@ -644,9 +644,16 @@ func (s *Store) ClearContextItems(ctx context.Context, convID int64) error {
// DeleteMessagesAfterID deletes all messages with ID > afterID for a conversation. // DeleteMessagesAfterID deletes all messages with ID > afterID for a conversation.
// Also clears related context_items, message_parts, summary_messages, and FTS entries. // Also clears related context_items, message_parts, summary_messages, and FTS entries.
// Uses transaction to ensure atomicity of the delete cascade.
func (s *Store) DeleteMessagesAfterID(ctx context.Context, convID int64, afterID int64) error { func (s *Store) DeleteMessagesAfterID(ctx context.Context, convID int64, afterID int64) error {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
// Get message IDs to delete for cleaning up related tables // Get message IDs to delete for cleaning up related tables
rows, err := s.db.QueryContext(ctx, rows, err := tx.QueryContext(ctx,
"SELECT message_id FROM messages WHERE conversation_id = ? AND message_id > ?", convID, afterID) "SELECT message_id FROM messages WHERE conversation_id = ? AND message_id > ?", convID, afterID)
if err != nil { if err != nil {
return err return err
@ -667,20 +674,29 @@ func (s *Store) DeleteMessagesAfterID(ctx context.Context, convID int64, afterID
// Delete context_items referencing these messages // Delete context_items referencing these messages
for _, msgID := range msgIDs { for _, msgID := range msgIDs {
s.db.ExecContext(ctx, "DELETE FROM context_items WHERE message_id = ?", msgID) if _, err := tx.ExecContext(ctx, "DELETE FROM context_items WHERE message_id = ?", msgID); err != nil {
return err
}
} }
// Delete from message_parts, summary_messages, and FTS // Delete from message_parts and summary_messages
// Note: messages_fts is handled automatically by trigger, no manual delete needed
for _, msgID := range msgIDs { for _, msgID := range msgIDs {
s.db.ExecContext(ctx, "DELETE FROM message_parts WHERE message_id = ?", msgID) if _, err := tx.ExecContext(ctx, "DELETE FROM message_parts WHERE message_id = ?", msgID); err != nil {
s.db.ExecContext(ctx, "DELETE FROM summary_messages WHERE message_id = ?", msgID) return err
s.db.ExecContext(ctx, "DELETE FROM messages_fts WHERE message_id = ?", msgID) }
if _, err := tx.ExecContext(ctx, "DELETE FROM summary_messages WHERE message_id = ?", msgID); err != nil {
return err
}
} }
// Delete messages // Delete messages
_, err = s.db.ExecContext(ctx, if _, err := tx.ExecContext(ctx,
"DELETE FROM messages WHERE conversation_id = ? AND message_id > ?", convID, afterID) "DELETE FROM messages WHERE conversation_id = ? AND message_id > ?", convID, afterID); err != nil {
return err return err
}
return tx.Commit()
} }
// AppendContextMessage appends a single message to context_items at next ordinal. // AppendContextMessage appends a single message to context_items at next ordinal.
@ -707,7 +723,13 @@ func (s *Store) AppendContextSummary(ctx context.Context, convID int64, summaryI
} }
func (s *Store) appendContextItems(ctx context.Context, convID int64, items []ContextItem) error { func (s *Store) appendContextItems(ctx context.Context, convID int64, items []ContextItem) error {
maxOrd, err := s.GetMaxOrdinal(ctx, convID) tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
maxOrd, err := s.GetMaxOrdinalTx(ctx, tx, convID)
if err != nil { if err != nil {
return err return err
} }
@ -720,10 +742,10 @@ func (s *Store) appendContextItems(ctx context.Context, convID int64, items []Co
// Resolve token count if not set // Resolve token count if not set
tokenCount := item.TokenCount tokenCount := item.TokenCount
if tokenCount == 0 { if tokenCount == 0 {
tokenCount = s.resolveItemTokenCount(ctx, item) tokenCount = s.resolveItemTokenCountTx(ctx, tx, item)
} }
_, err = s.db.ExecContext(ctx, _, err = tx.ExecContext(ctx,
`INSERT INTO context_items (conversation_id, ordinal, item_type, summary_id, message_id, token_count) `INSERT INTO context_items (conversation_id, ordinal, item_type, summary_id, message_id, token_count)
VALUES (?, ?, ?, ?, ?, ?)`, VALUES (?, ?, ?, ?, ?, ?)`,
convID, ordinal, item.ItemType, convID, ordinal, item.ItemType,
@ -735,14 +757,14 @@ func (s *Store) appendContextItems(ctx context.Context, convID int64, items []Co
} }
ordinal += OrdinalStep ordinal += OrdinalStep
} }
return nil return tx.Commit()
} }
// resolveItemTokenCount looks up token count from message or summary if not provided. // resolveItemTokenCountTx looks up token count within a transaction.
func (s *Store) resolveItemTokenCount(ctx context.Context, item ContextItem) int { func (s *Store) resolveItemTokenCountTx(ctx context.Context, tx *sql.Tx, item ContextItem) int {
if item.ItemType == "message" && item.MessageID > 0 { if item.ItemType == "message" && item.MessageID > 0 {
var tc int var tc int
err := s.db.QueryRowContext(ctx, err := tx.QueryRowContext(ctx,
"SELECT token_count FROM messages WHERE message_id = ?", item.MessageID, "SELECT token_count FROM messages WHERE message_id = ?", item.MessageID,
).Scan(&tc) ).Scan(&tc)
if err == nil { if err == nil {
@ -751,7 +773,7 @@ func (s *Store) resolveItemTokenCount(ctx context.Context, item ContextItem) int
} }
if item.ItemType == "summary" && item.SummaryID != "" { if item.ItemType == "summary" && item.SummaryID != "" {
var tc int var tc int
err := s.db.QueryRowContext(ctx, err := tx.QueryRowContext(ctx,
"SELECT token_count FROM summaries WHERE summary_id = ?", item.SummaryID, "SELECT token_count FROM summaries WHERE summary_id = ?", item.SummaryID,
).Scan(&tc) ).Scan(&tc)
if err == nil { if err == nil {
@ -1028,6 +1050,22 @@ func (s *Store) GetMaxOrdinal(ctx context.Context, convID int64) (int, error) {
return int(maxOrd.Int64), nil return int(maxOrd.Int64), nil
} }
// GetMaxOrdinalTx returns the highest ordinal within a transaction.
func (s *Store) GetMaxOrdinalTx(ctx context.Context, tx *sql.Tx, convID int64) (int, error) {
var maxOrd sql.NullInt64
err := tx.QueryRowContext(ctx,
"SELECT MAX(ordinal) FROM context_items WHERE conversation_id = ?",
convID,
).Scan(&maxOrd)
if err != nil {
return 0, err
}
if !maxOrd.Valid {
return 0, nil
}
return int(maxOrd.Int64), nil
}
// GetDistinctDepthsInContext returns distinct depth levels of summaries currently in context. // GetDistinctDepthsInContext returns distinct depth levels of summaries currently in context.
// maxOrdinalExclusive filters out summaries with ordinal >= this value (0 = no filter). // maxOrdinalExclusive filters out summaries with ordinal >= this value (0 = no filter).
func (s *Store) GetDistinctDepthsInContext(ctx context.Context, convID int64, maxOrdinalExclusive int) ([]int, error) { func (s *Store) GetDistinctDepthsInContext(ctx context.Context, convID int64, maxOrdinalExclusive int) ([]int, error) {