diff --git a/internal/opencode/grep.go b/internal/opencode/grep.go index 5b9a384..5d0850d 100644 --- a/internal/opencode/grep.go +++ b/internal/opencode/grep.go @@ -1,9 +1,11 @@ package opencode import ( + "fmt" "strings" "github.com/sorafujitani/ccsession/internal/grep" + "github.com/sorafujitani/ccsession/internal/session" ) // GrepKeys returns the set of root session ids whose title or message text @@ -30,14 +32,16 @@ func (d *DB) GrepKeys(query string, regex bool) (map[string]struct{}, error) { } set := make(map[string]struct{}) + var bodyRoots []rootRow for _, r := range roots { - hit, err := d.sessionMatches(r, match) - if err != nil { - return nil, err - } - if hit { + if match(r.title) { set[r.id] = struct{}{} + continue } + bodyRoots = append(bodyRoots, r) + } + if err := d.matchBodies(bodyRoots, match, set); err != nil { + return nil, err } return set, nil } @@ -84,20 +88,162 @@ WHERE parent_id IS NULL AND time_archived IS NULL` return out, rows.Err() } -func (d *DB) sessionMatches(r rootRow, match func(string) bool) (bool, error) { - if match(r.title) { - return true, nil +func (d *DB) matchBodies(roots []rootRow, match func(string) bool, set map[string]struct{}) error { + if len(roots) == 0 { + return nil + } + for _, chunk := range chunkRoots(roots, 500) { + hasProjection, err := d.matchProjectionBodies(chunk, match, set) + if err != nil { + return err + } + var fallback []rootRow + for _, r := range chunk { + if _, ok := hasProjection[r.id]; !ok { + fallback = append(fallback, r) + } + } + if err := d.matchPartBodies(fallback, match, set); err != nil { + return err + } + } + return nil +} + +func (d *DB) matchProjectionBodies(roots []rootRow, match func(string) bool, set map[string]struct{}) (map[string]struct{}, error) { + hasProjection := make(map[string]struct{}) + if len(roots) == 0 { + return hasProjection, nil } - msgs, _, _, err := d.Messages(r.id, 0) + q, args := inQuery(`SELECT session_id, type, data +FROM session_message +WHERE session_id IN (%s) AND type IN ('user', 'assistant') +ORDER BY session_id, seq, id`, rootIDs(roots)) + rows, err := d.query(q, args...) if err != nil { - return false, err + if isMissingTable(err) { + return hasProjection, nil + } + return nil, err + } + defer rows.Close() + for rows.Next() { + var sessionID string + var typ string + var data []byte + if err := rows.Scan(&sessionID, &typ, &data); err != nil { + return nil, err + } + msg := projectionMessage(typ, data) + if msg.Body == "" { + continue + } + hasProjection[sessionID] = struct{}{} + if _, matched := set[sessionID]; matched { + continue + } + if match(msg.Body) { + set[sessionID] = struct{}{} + } } - for _, m := range msgs { - if match(m.Body) { - return true, nil + return hasProjection, rows.Err() +} + +func (d *DB) matchPartBodies(roots []rootRow, match func(string) bool, set map[string]struct{}) error { + if len(roots) == 0 { + return nil + } + q, args := inQuery(`SELECT m.session_id, m.id, m.data, p.data +FROM message m +LEFT JOIN part p ON p.message_id = m.id +WHERE m.session_id IN (%s) +ORDER BY m.session_id, m.time_created, m.id, p.id`, rootIDs(roots)) + rows, err := d.query(q, args...) + if err != nil { + return err + } + defer rows.Close() + + var ( + curSession string + curMsgID string + cur *sessionBody + ) + flush := func() { + if cur == nil || !isRenderableRole(cur.role) { + return + } + if _, matched := set[curSession]; matched { + return + } + if match(cur.body) { + set[curSession] = struct{}{} + } + } + for rows.Next() { + var sessionID string + var msgID string + var msgData []byte + var part []byte + if err := rows.Scan(&sessionID, &msgID, &msgData, &part); err != nil { + return err + } + if _, matched := set[sessionID]; matched { + continue + } + if cur == nil || sessionID != curSession || msgID != curMsgID { + flush() + curSession = sessionID + curMsgID = msgID + msg := newTurn(msgData) + cur = &sessionBody{role: msg.Role} + } + appendBodyText(&cur.body, part) + } + flush() + return rows.Err() +} + +type sessionBody struct { + role string + body string +} + +func appendBodyText(body *string, partData []byte) { + msg := &session.Message{Body: *body} + appendText(msg, partData) + *body = msg.Body +} + +func chunkRoots(roots []rootRow, size int) [][]rootRow { + var chunks [][]rootRow + for len(roots) > 0 { + n := min(len(roots), size) + chunks = append(chunks, roots[:n]) + roots = roots[n:] + } + return chunks +} + +func rootIDs(roots []rootRow) []string { + ids := make([]string, len(roots)) + for i, r := range roots { + ids[i] = r.id + } + return ids +} + +func inQuery(format string, ids []string) (string, []any) { + var b strings.Builder + args := make([]any, len(ids)) + for i, id := range ids { + if i > 0 { + b.WriteByte(',') } + b.WriteByte('?') + args[i] = id } - return false, nil + return fmt.Sprintf(format, b.String()), args } // asciiFoldTargets are the ASCII letters a non-ASCII rune lower-cases onto diff --git a/internal/opencode/grep_test.go b/internal/opencode/grep_test.go index 7a162c3..c69556c 100644 --- a/internal/opencode/grep_test.go +++ b/internal/opencode/grep_test.go @@ -151,6 +151,36 @@ func TestGrepKeys_FindsUnicodeCaseFoldMatch(t *testing.T) { } } +func TestGrepKeys_BatchedBodiesWhenPrefilterDisabled(t *testing.T) { + f := newFixture(t, fixtureOpts{}) + for i := range 40 { + id := "ses_batch_" + itoa(int64(i)) + f.session(id, "/p/"+id, "zzz", int64(i+1)*100) + body := "zzz" + if i == 37 { + body = "İstanbul notes" + } + f.partsTurn(id, "user", int64(i+1)*100, body) + } + + got := keysSorted(mustGrep(t, f.open(), "i", false)) + if !slices.Equal(got, []string{"ses_batch_37"}) { + t.Fatalf("batched disabled-prefilter grep = %v, want [ses_batch_37]", got) + } +} + +func TestGrepKeys_ProjectionPreferredOverParts(t *testing.T) { + f := newFixture(t, fixtureOpts{}) + f.session("ses_projection", "/p", "ordinary", 100) + f.projectionRow("ses_projection", "user", 1, `{"text":"projection text","time":{"created":10}}`) + f.partsTurn("ses_projection", "user", 10, "needle only in ignored parts") + + got := mustGrep(t, f.open(), "needle", false) + if contains(got, "ses_projection") { + t.Fatalf("grep matched parts despite renderable projection: %v", keysSorted(got)) + } +} + func contains(m map[string]struct{}, k string) bool { _, ok := m[k] return ok @@ -228,3 +258,37 @@ func BenchmarkGrepKeysHitAndMiss(b *testing.B) { }) } } + +func BenchmarkGrepKeysPrefilterDisabled(b *testing.B) { + f := newFixture(b, fixtureOpts{}) + for i := range 512 { + id := "ses_disabled_" + itoa(int64(i)) + f.session(id, "/tmp/"+id, "zzz", int64(i+1)*1000) + body := "zzz" + if i%64 == 0 { + body = "İstanbul notes" + } + f.partsTurn(id, "user", int64(i+1)*1000, body) + } + d := f.open() + + for _, tc := range []struct { + name string + query string + }{ + {name: "hit", query: "i"}, + {name: "miss", query: "k"}, + } { + b.Run(tc.name, func(b *testing.B) { + for range b.N { + keys, err := d.GrepKeys(tc.query, false) + if err != nil { + b.Fatalf("GrepKeys: %v", err) + } + if keys == nil { + b.Fatal("GrepKeys returned nil set") + } + } + }) + } +} diff --git a/internal/opencode/messages.go b/internal/opencode/messages.go index e8434a0..b0e44f2 100644 --- a/internal/opencode/messages.go +++ b/internal/opencode/messages.go @@ -1,6 +1,7 @@ package opencode import ( + "database/sql" "encoding/json" "slices" "strings" @@ -16,24 +17,16 @@ import ( // The session_message projection is read first; it falls back to the message // +part store when the projection is empty of renderable turns or absent. func (d *DB) Messages(sessionID string, limit int) (msgs []session.Message, startedAt time.Time, total int, err error) { - all, err := d.messagesFromProjection(sessionID) + all, startedAt, total, err := d.messagesFromProjection(sessionID, limit) if err != nil { return nil, time.Time{}, 0, err } - if len(all) == 0 { - all, err = d.messagesFromParts(sessionID) + if total == 0 { + all, startedAt, total, err = d.messagesFromParts(sessionID, limit) if err != nil { return nil, time.Time{}, 0, err } } - if len(all) == 0 { - return nil, time.Time{}, 0, nil - } - startedAt = all[0].Timestamp - total = len(all) - if limit > 0 && len(all) > limit { - all = all[len(all)-limit:] - } return all, startedAt, total, nil } @@ -52,7 +45,10 @@ type textPart struct { // messagesFromParts groups each message's text parts into one turn, keeping // only user/assistant turns. The LEFT JOIN preserves a turn whose only parts // are non-text (reasoning/step-start) as an empty body rather than dropping it. -func (d *DB) messagesFromParts(sessionID string) ([]session.Message, error) { +func (d *DB) messagesFromParts(sessionID string, limit int) ([]session.Message, time.Time, int, error) { + if limit > 0 { + return d.limitedMessagesFromParts(sessionID, limit) + } const q = `SELECT m.id, m.data, p.data FROM message m LEFT JOIN part p ON p.message_id = m.id @@ -60,10 +56,70 @@ WHERE m.session_id = ? ORDER BY m.time_created, m.id, p.id` rows, err := d.query(q, sessionID) if err != nil { - return nil, err + return nil, time.Time{}, 0, err } defer rows.Close() + out, err := scanPartMessages(rows, nil) + if err != nil { + return nil, time.Time{}, 0, err + } + if len(out) == 0 { + return nil, time.Time{}, 0, nil + } + return out, out[0].Timestamp, len(out), nil +} + +func (d *DB) limitedMessagesFromParts(sessionID string, limit int) ([]session.Message, time.Time, int, error) { + total, startedAt, err := d.partsCountAndStartedAt(sessionID) + if err != nil || total == 0 { + return nil, time.Time{}, total, err + } + const q = `SELECT m.id, m.data, p.data +FROM ( + SELECT id, data, time_created + FROM message + WHERE session_id = ? AND (data LIKE '%"role":"user"%' OR data LIKE '%"role":"assistant"%') + ORDER BY time_created DESC, id DESC + LIMIT ? +) m +LEFT JOIN part p ON p.message_id = m.id +ORDER BY m.time_created, m.id, p.id` + rows, err := d.query(q, sessionID, limit) + if err != nil { + return nil, time.Time{}, 0, err + } + defer rows.Close() + out, err := scanPartMessages(rows, nil) + if err != nil { + return nil, time.Time{}, 0, err + } + return out, startedAt, total, nil +} + +func (d *DB) partsCountAndStartedAt(sessionID string) (int, time.Time, error) { + const q = `SELECT COUNT(*), COALESCE(MIN(time_created), 0) +FROM message +WHERE session_id = ? AND (data LIKE '%"role":"user"%' OR data LIKE '%"role":"assistant"%')` + rows, err := d.query(q, sessionID) + if err != nil { + return 0, time.Time{}, err + } + defer rows.Close() + var total int + var startedMs int64 + if rows.Next() { + if err := rows.Scan(&total, &startedMs); err != nil { + return 0, time.Time{}, err + } + } + if err := rows.Err(); err != nil { + return 0, time.Time{}, err + } + return total, msToTime(startedMs), nil +} + +func scanPartMessages(rows *sql.Rows, sessionIDs map[string]struct{}) ([]session.Message, error) { var ( out []session.Message curID string @@ -78,8 +134,16 @@ ORDER BY m.time_created, m.id, p.id` var msgID string var msgData []byte var part []byte // NULL for a message with no parts - if err := rows.Scan(&msgID, &msgData, &part); err != nil { - return nil, err + if sessionIDs == nil { + if err := rows.Scan(&msgID, &msgData, &part); err != nil { + return nil, err + } + } else { + var sessionID string + if err := rows.Scan(&sessionID, &msgID, &msgData, &part); err != nil { + return nil, err + } + sessionIDs[sessionID] = struct{}{} } if cur == nil || msgID != curID { flush() @@ -123,7 +187,10 @@ type projectionData struct { // chronological order; type comes from the column, not the JSON. A projection // empty of renderable turns, or absent (a DB predating it), reads as nil so the // parts store takes over. -func (d *DB) messagesFromProjection(sessionID string) ([]session.Message, error) { +func (d *DB) messagesFromProjection(sessionID string, limit int) ([]session.Message, time.Time, int, error) { + if limit > 0 { + return d.limitedMessagesFromProjection(sessionID, limit) + } const q = `SELECT type, data FROM session_message WHERE session_id = ? @@ -131,9 +198,9 @@ ORDER BY seq DESC, id DESC` rows, err := d.query(q, sessionID) if err != nil { if isMissingTable(err) { - return nil, nil + return nil, time.Time{}, 0, nil } - return nil, err + return nil, time.Time{}, 0, err } defer rows.Close() @@ -142,27 +209,98 @@ ORDER BY seq DESC, id DESC` var typ string var data []byte if err := rows.Scan(&typ, &data); err != nil { - return nil, err + return nil, time.Time{}, 0, err } if !isRenderableRole(typ) { continue } - var pd projectionData - _ = json.Unmarshal(data, &pd) - if pd.Text == "" { + msg := projectionMessage(typ, data) + if msg.Body == "" { continue } - rev = append(rev, session.Message{ - Role: typ, - Timestamp: msToTime(pd.Time.Created), - Body: pd.Text, - }) + rev = append(rev, msg) + } + if err := rows.Err(); err != nil { + return nil, time.Time{}, 0, err + } + slices.Reverse(rev) + if len(rev) == 0 { + return nil, time.Time{}, 0, nil + } + return rev, rev[0].Timestamp, len(rev), nil +} + +func (d *DB) limitedMessagesFromProjection(sessionID string, limit int) ([]session.Message, time.Time, int, error) { + total, startedAt, err := d.projectionCountAndStartedAt(sessionID) + if err != nil || total == 0 { + return nil, time.Time{}, total, err + } + const q = `SELECT type, data +FROM session_message +WHERE session_id = ? + AND type IN ('user', 'assistant') + AND COALESCE(json_extract(data, '$.text'), '') != '' +ORDER BY seq DESC, id DESC +LIMIT ?` + rows, err := d.query(q, sessionID, limit) + if err != nil { + if isMissingTable(err) { + return nil, time.Time{}, 0, nil + } + return nil, time.Time{}, 0, err + } + defer rows.Close() + var rev []session.Message + for rows.Next() { + var typ string + var data []byte + if err := rows.Scan(&typ, &data); err != nil { + return nil, time.Time{}, 0, err + } + rev = append(rev, projectionMessage(typ, data)) } if err := rows.Err(); err != nil { - return nil, err + return nil, time.Time{}, 0, err } slices.Reverse(rev) - return rev, nil + return rev, startedAt, total, nil +} + +func (d *DB) projectionCountAndStartedAt(sessionID string) (int, time.Time, error) { + const q = `SELECT COUNT(*), COALESCE(MIN(json_extract(data, '$.time.created')), 0) +FROM session_message +WHERE session_id = ? + AND type IN ('user', 'assistant') + AND COALESCE(json_extract(data, '$.text'), '') != ''` + rows, err := d.query(q, sessionID) + if err != nil { + if isMissingTable(err) { + return 0, time.Time{}, nil + } + return 0, time.Time{}, err + } + defer rows.Close() + var total int + var startedMs int64 + if rows.Next() { + if err := rows.Scan(&total, &startedMs); err != nil { + return 0, time.Time{}, err + } + } + if err := rows.Err(); err != nil { + return 0, time.Time{}, err + } + return total, msToTime(startedMs), nil +} + +func projectionMessage(typ string, data []byte) session.Message { + var pd projectionData + _ = json.Unmarshal(data, &pd) + return session.Message{ + Role: typ, + Timestamp: msToTime(pd.Time.Created), + Body: pd.Text, + } } func isRenderableRole(role string) bool { diff --git a/internal/opencode/messages_test.go b/internal/opencode/messages_test.go index a3eda26..695b832 100644 --- a/internal/opencode/messages_test.go +++ b/internal/opencode/messages_test.go @@ -169,6 +169,25 @@ func TestMessages_LimitReturnsNewestN(t *testing.T) { } } +func TestMessages_LimitEqualCountReturnsAll(t *testing.T) { + f := newFixture(t, fixtureOpts{}) + f.session("ses_a", "/p", "t", 100) + for i := 1; i <= 3; i++ { + f.partsTurn("ses_a", "user", int64(i*10), itoa(int64(i))) + } + + msgs, started, total, err := f.open().Messages("ses_a", 3) + if err != nil { + t.Fatal(err) + } + if total != 3 || started.UnixMilli() != 10 { + t.Fatalf("total/started = %d/%d, want 3/10", total, started.UnixMilli()) + } + if got, want := bodies(msgs), []string{"user:1", "user:2", "user:3"}; !slices.Equal(got, want) { + t.Fatalf("limited equal messages = %v, want %v", got, want) + } +} + func TestMessages_NonTextPartCountsAsEmptyTurn(t *testing.T) { f := newFixture(t, fixtureOpts{}) f.session("ses_a", "/p", "t", 100)