mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: Add full text search over chat messages (#27126)
Closes CODAGT-721 Closes CODAGT-722 Closes CODAGT-723 Closes CODAGT-724 Closes CODAGT-725 This PR adds the database and API pieces necessary to support full-text chat message search. - Adds required chat schema for full-text search - Adds dbpurge job to populate search_tsv in the background - Adds `search` parameter to GetChats query - Adds `search` filter to `searchquery.Chats` - Wires chat search filter into chats API > Implemented by Coder Agents, reviewed and tested by a human.
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
-- Restore the original trigger bodies from 000519.
|
||||
CREATE OR REPLACE FUNCTION set_chat_message_revision_before()
|
||||
RETURNS trigger AS $$
|
||||
DECLARE
|
||||
chat_snapshot_version bigint;
|
||||
BEGIN
|
||||
IF TG_OP = 'INSERT' AND NEW.revision IS NOT NULL THEN
|
||||
RAISE EXCEPTION 'chat_messages.revision must be assigned by trigger';
|
||||
END IF;
|
||||
|
||||
IF TG_OP = 'UPDATE' THEN
|
||||
IF OLD.chat_id IS DISTINCT FROM NEW.chat_id THEN
|
||||
RAISE EXCEPTION 'chat_messages.chat_id is immutable';
|
||||
END IF;
|
||||
|
||||
IF OLD.revision IS DISTINCT FROM NEW.revision THEN
|
||||
RAISE EXCEPTION 'chat_messages.revision must be assigned by trigger';
|
||||
END IF;
|
||||
|
||||
IF OLD IS NOT DISTINCT FROM NEW THEN
|
||||
RETURN NEW;
|
||||
END IF;
|
||||
END IF;
|
||||
|
||||
SELECT snapshot_version INTO chat_snapshot_version
|
||||
FROM chats WHERE id = NEW.chat_id;
|
||||
|
||||
IF chat_snapshot_version IS NULL THEN
|
||||
RAISE EXCEPTION 'chat % does not exist', NEW.chat_id;
|
||||
END IF;
|
||||
|
||||
NEW.revision = chat_snapshot_version;
|
||||
RETURN NEW;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
CREATE OR REPLACE FUNCTION update_chat_history_after_message_update()
|
||||
RETURNS trigger AS $$
|
||||
BEGIN
|
||||
UPDATE chats c
|
||||
SET history_version = c.snapshot_version,
|
||||
generation_attempt = 0
|
||||
FROM (
|
||||
SELECT DISTINCT n.chat_id
|
||||
FROM chat_message_history_new_rows n
|
||||
JOIN chat_message_history_old_rows o ON o.id = n.id
|
||||
WHERE o IS DISTINCT FROM n
|
||||
) AS affected
|
||||
WHERE c.id = affected.chat_id
|
||||
AND (
|
||||
c.history_version IS DISTINCT FROM c.snapshot_version
|
||||
OR c.generation_attempt <> 0
|
||||
);
|
||||
RETURN NULL;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
DROP INDEX IF EXISTS idx_chat_diff_statuses_pr_title_fts;
|
||||
|
||||
DROP INDEX IF EXISTS idx_chats_title_fts;
|
||||
|
||||
DROP INDEX IF EXISTS idx_chat_messages_search_tsv_pending;
|
||||
|
||||
DROP INDEX IF EXISTS idx_chat_messages_search_tsv;
|
||||
|
||||
ALTER TABLE chat_messages DROP COLUMN IF EXISTS search_tsv;
|
||||
|
||||
DROP FUNCTION IF EXISTS chat_message_search_text(jsonb);
|
||||
@@ -0,0 +1,97 @@
|
||||
CREATE FUNCTION chat_message_search_text(content jsonb) RETURNS text
|
||||
LANGUAGE sql IMMUTABLE PARALLEL SAFE AS $$
|
||||
SELECT CASE WHEN jsonb_typeof(content) = 'array' THEN (
|
||||
SELECT string_agg(part->>'text', ' ' ORDER BY ordinality)
|
||||
FROM jsonb_array_elements(content) WITH ORDINALITY AS t(part, ordinality)
|
||||
WHERE part->>'type' = 'text'
|
||||
) END
|
||||
$$;
|
||||
|
||||
COMMENT ON FUNCTION chat_message_search_text IS 'Extracts searchable content from chat_messages. Returns NULL for scalar JSON strings (content_version=0). Immutable as it is used in indexes.';
|
||||
|
||||
-- Populated by a background sweep, not at insert time. NULL means pending.
|
||||
ALTER TABLE chat_messages ADD COLUMN search_tsv tsvector;
|
||||
|
||||
COMMENT ON COLUMN chat_messages.search_tsv IS 'Used for full text search. NULL initially, populated async via background job.';
|
||||
|
||||
CREATE INDEX idx_chat_messages_search_tsv ON chat_messages
|
||||
USING GIN (search_tsv)
|
||||
WHERE ((search_tsv IS NOT NULL) AND (deleted = false) AND (visibility = ANY (ARRAY['user'::chat_message_visibility, 'both'::chat_message_visibility])) AND (role = ANY (ARRAY['user'::chat_message_role, 'assistant'::chat_message_role])));
|
||||
|
||||
COMMENT ON INDEX idx_chat_messages_search_tsv IS 'Partial index over chat_messages used for full text search. Only defined over ''searchable'' rows of chat_messages.';
|
||||
|
||||
CREATE INDEX idx_chat_messages_search_tsv_pending ON chat_messages USING btree (id DESC)
|
||||
WHERE ((search_tsv IS NULL) AND (deleted = false) AND (visibility = ANY (ARRAY['user'::chat_message_visibility, 'both'::chat_message_visibility])) AND (role = ANY (ARRAY['user'::chat_message_role, 'assistant'::chat_message_role])));
|
||||
|
||||
COMMENT ON INDEX idx_chat_messages_search_tsv IS 'Partial index over chat_messages used for populating search_tsv in the background. Only defined over ''searchable'' rows of chat_messages where search_tsv is NULL.';
|
||||
|
||||
CREATE INDEX idx_chats_title_fts ON chats USING GIN (to_tsvector('simple', title));
|
||||
|
||||
COMMENT ON index idx_chats_title_fts IS 'Used for full text search. Defined over all rows of the chats table.';
|
||||
|
||||
CREATE INDEX idx_chat_diff_statuses_pr_title_fts ON chat_diff_statuses USING GIN (to_tsvector('simple', pull_request_title));
|
||||
|
||||
COMMENT ON index idx_chats_title_fts IS 'Used for full text search. Defined over all rows of the chats table.';
|
||||
|
||||
CREATE OR REPLACE FUNCTION set_chat_message_revision_before()
|
||||
RETURNS trigger AS $$
|
||||
DECLARE
|
||||
chat_snapshot_version bigint;
|
||||
cmp chat_messages;
|
||||
BEGIN
|
||||
IF TG_OP = 'INSERT' AND NEW.revision IS NOT NULL THEN
|
||||
RAISE EXCEPTION 'chat_messages.revision must be assigned by trigger';
|
||||
END IF;
|
||||
|
||||
IF TG_OP = 'UPDATE' THEN
|
||||
IF OLD.chat_id IS DISTINCT FROM NEW.chat_id THEN
|
||||
RAISE EXCEPTION 'chat_messages.chat_id is immutable';
|
||||
END IF;
|
||||
|
||||
IF OLD.revision IS DISTINCT FROM NEW.revision THEN
|
||||
RAISE EXCEPTION 'chat_messages.revision must be assigned by trigger';
|
||||
END IF;
|
||||
|
||||
cmp := NEW;
|
||||
cmp.search_tsv := OLD.search_tsv;
|
||||
IF OLD IS NOT DISTINCT FROM cmp THEN
|
||||
RETURN NEW;
|
||||
END IF;
|
||||
END IF;
|
||||
|
||||
SELECT snapshot_version INTO chat_snapshot_version
|
||||
FROM chats WHERE id = NEW.chat_id;
|
||||
|
||||
IF chat_snapshot_version IS NULL THEN
|
||||
RAISE EXCEPTION 'chat % does not exist', NEW.chat_id;
|
||||
END IF;
|
||||
|
||||
NEW.revision = chat_snapshot_version;
|
||||
RETURN NEW;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
COMMENT ON FUNCTION set_chat_message_revision_before IS 'Component of chatd. Updates chat_snapshot_version when any fields of chat_messages change. Excludes changes to search_tsv as it is not relevant to chatd''s processing loop.';
|
||||
|
||||
CREATE OR REPLACE FUNCTION update_chat_history_after_message_update()
|
||||
RETURNS trigger AS $$
|
||||
BEGIN
|
||||
UPDATE chats c
|
||||
SET history_version = c.snapshot_version,
|
||||
generation_attempt = 0
|
||||
FROM (
|
||||
SELECT DISTINCT n.chat_id
|
||||
FROM chat_message_history_new_rows n
|
||||
JOIN chat_message_history_old_rows o ON o.id = n.id
|
||||
WHERE (to_jsonb(o) - 'search_tsv') IS DISTINCT FROM (to_jsonb(n) - 'search_tsv')
|
||||
) AS affected
|
||||
WHERE c.id = affected.chat_id
|
||||
AND (
|
||||
c.history_version IS DISTINCT FROM c.snapshot_version
|
||||
OR c.generation_attempt <> 0
|
||||
);
|
||||
RETURN NULL;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
COMMENT ON FUNCTION update_chat_history_after_message_update IS 'Component of chatd. Updates history_version and generation_attempt on chats when chat_messages is updated. Excludes changes to search_tsv.';
|
||||
@@ -19,11 +19,13 @@ import (
|
||||
"github.com/golang-migrate/migrate/v4/source/stub"
|
||||
"github.com/google/uuid"
|
||||
"github.com/lib/pq"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
"golang.org/x/sync/errgroup"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/database/migrations"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -1866,3 +1868,244 @@ func TestMigration000498SoftDeleteStaleWorkspaceAgents(t *testing.T) {
|
||||
// TestSoftDeleteWorkspaceAgentsByWorkspaceID, plus integration tests
|
||||
// under coderd/coderd_test.go; not retested here.
|
||||
}
|
||||
|
||||
func TestMigration000543ChatMessageSearchText(t *testing.T) {
|
||||
t.Parallel()
|
||||
if testing.Short() {
|
||||
t.SkipNow()
|
||||
}
|
||||
|
||||
sqlDB := testSQLDB(t)
|
||||
require.NoError(t, migrations.Up(sqlDB))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
content sql.NullString
|
||||
want sql.NullString
|
||||
}{
|
||||
{
|
||||
name: "SingleTextPart",
|
||||
content: sql.NullString{String: `[{"type":"text","text":"hello world"}]`, Valid: true},
|
||||
want: sql.NullString{String: "hello world", Valid: true},
|
||||
},
|
||||
{
|
||||
name: "TextInterleavedWithNonText",
|
||||
content: sql.NullString{String: `[
|
||||
{"type":"text","text":"first"},
|
||||
{"type":"reasoning","text":"thinking"},
|
||||
{"type":"tool-call","toolName":"execute"},
|
||||
{"type":"text","text":"second"}
|
||||
]`, Valid: true},
|
||||
want: sql.NullString{String: "first second", Valid: true},
|
||||
},
|
||||
{
|
||||
name: "OnlyNonTextParts",
|
||||
content: sql.NullString{String: `[{"type":"reasoning","text":"thinking"}]`, Valid: true},
|
||||
want: sql.NullString{},
|
||||
},
|
||||
{
|
||||
name: "ScalarContent",
|
||||
content: sql.NullString{String: `"hello"`, Valid: true},
|
||||
want: sql.NullString{},
|
||||
},
|
||||
{
|
||||
name: "EmptyArray",
|
||||
content: sql.NullString{String: `[]`, Valid: true},
|
||||
want: sql.NullString{},
|
||||
},
|
||||
{
|
||||
name: "NullInput",
|
||||
content: sql.NullString{},
|
||||
want: sql.NullString{},
|
||||
},
|
||||
{
|
||||
name: "ElementsMissingTypeOrText",
|
||||
content: sql.NullString{String: `[{"text":"no type"},{"type":"text"},{"type":"text","text":"kept"}]`, Valid: true},
|
||||
want: sql.NullString{String: "kept", Valid: true},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
var got sql.NullString
|
||||
err := sqlDB.QueryRowContext(ctx,
|
||||
`SELECT chat_message_search_text($1::jsonb)`, tc.content,
|
||||
).Scan(&got)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tc.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Shared eligibility predicate of the two partial chat_messages search
|
||||
// indexes. Queries must repeat it verbatim.
|
||||
const eligibilityPredicate = `deleted = false
|
||||
AND visibility IN ('user', 'both')
|
||||
AND role IN ('user', 'assistant')`
|
||||
|
||||
func TestMigration000543ChatSearchSchemaIndexes(t *testing.T) {
|
||||
t.Parallel()
|
||||
if testing.Short() {
|
||||
t.SkipNow()
|
||||
}
|
||||
|
||||
sqlDB := testSQLDB(t)
|
||||
require.NoError(t, migrations.Up(sqlDB))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
table string
|
||||
partial bool
|
||||
}{
|
||||
{name: "idx_chat_messages_search_tsv", table: "chat_messages", partial: true},
|
||||
{name: "idx_chat_messages_search_tsv_pending", table: "chat_messages", partial: true},
|
||||
{name: "idx_chats_title_fts", table: "chats", partial: false},
|
||||
{name: "idx_chat_diff_statuses_pr_title_fts", table: "chat_diff_statuses", partial: false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
var table string
|
||||
var partial bool
|
||||
err := sqlDB.QueryRowContext(ctx, `
|
||||
SELECT i.tablename, x.indpred IS NOT NULL
|
||||
FROM pg_indexes i
|
||||
JOIN pg_class c ON c.relname = i.indexname
|
||||
JOIN pg_index x ON x.indexrelid = c.oid
|
||||
WHERE i.indexname = $1`, tc.name,
|
||||
).Scan(&table, &partial)
|
||||
require.NoError(t, err, "index %s should exist", tc.name)
|
||||
require.Equal(t, tc.table, table, "index %s table", tc.name)
|
||||
require.Equal(t, tc.partial, partial, "index %s partial", tc.name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigration000543ChatSearchSchemaBehavior(t *testing.T) {
|
||||
t.Parallel()
|
||||
if testing.Short() {
|
||||
t.SkipNow()
|
||||
}
|
||||
|
||||
sqlDB := testSQLDB(t)
|
||||
require.NoError(t, migrations.Up(sqlDB))
|
||||
db := database.New(sqlDB)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
owner := dbgen.User(t, db, database.User{})
|
||||
_ = dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"})
|
||||
modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
IsDefault: true,
|
||||
})
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: owner.ID,
|
||||
LastModelConfigID: modelCfg.ID,
|
||||
})
|
||||
|
||||
newMsg := func(role database.ChatMessageRole, visibility database.ChatMessageVisibility, content string) database.ChatMessage {
|
||||
seed := database.ChatMessage{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true},
|
||||
Role: role,
|
||||
Visibility: visibility,
|
||||
}
|
||||
if content != "" {
|
||||
seed.Content = pqtype.NullRawMessage{RawMessage: []byte(content), Valid: true}
|
||||
}
|
||||
return dbgen.ChatMessage(t, db, seed)
|
||||
}
|
||||
textContent := func(text string) string {
|
||||
return `[{"type":"text","text":"` + text + `"}]`
|
||||
}
|
||||
|
||||
pendingIDs := func(ctx context.Context, limit int) []int64 {
|
||||
rows, err := sqlDB.QueryContext(ctx, `
|
||||
SELECT id FROM chat_messages
|
||||
WHERE search_tsv IS NULL AND `+eligibilityPredicate+`
|
||||
ORDER BY id DESC
|
||||
LIMIT $1`, limit)
|
||||
require.NoError(t, err)
|
||||
defer rows.Close()
|
||||
var ids []int64
|
||||
for rows.Next() {
|
||||
var id int64
|
||||
require.NoError(t, rows.Scan(&id))
|
||||
ids = append(ids, id)
|
||||
}
|
||||
require.NoError(t, rows.Err())
|
||||
return ids
|
||||
}
|
||||
|
||||
// Insert regression: RETURNING * must survive the new column, and new
|
||||
// rows must start with search_tsv NULL so they enter the pending queue.
|
||||
eligibleText := newMsg(database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, textContent("deploy the search feature"))
|
||||
var tsvIsNull bool
|
||||
err := sqlDB.QueryRowContext(ctx,
|
||||
`SELECT search_tsv IS NULL FROM chat_messages WHERE id = $1`, eligibleText.ID,
|
||||
).Scan(&tsvIsNull)
|
||||
require.NoError(t, err)
|
||||
require.True(t, tsvIsNull, "new rows must have search_tsv NULL")
|
||||
|
||||
eligibleNoText := newMsg(database.ChatMessageRoleAssistant, database.ChatMessageVisibilityUser, `[{"type":"reasoning","text":"thinking"}]`)
|
||||
toolMsg := newMsg(database.ChatMessageRoleTool, database.ChatMessageVisibilityBoth, textContent("tool output about deploy"))
|
||||
modelOnly := newMsg(database.ChatMessageRoleUser, database.ChatMessageVisibilityModel, textContent("model-only deploy note"))
|
||||
deletedMsg := newMsg(database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, textContent("deleted deploy message"))
|
||||
_, err = sqlDB.ExecContext(ctx, `UPDATE chat_messages SET deleted = true WHERE id = $1`, deletedMsg.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Only eligible rows appear in the queue, newest first. The tool-role,
|
||||
// model-only, and soft-deleted rows are excluded even though their
|
||||
// search_tsv is NULL.
|
||||
require.Equal(t, []int64{eligibleNoText.ID, eligibleText.ID}, pendingIDs(ctx, 10))
|
||||
|
||||
// Sweep-style UPDATE. The '' sentinel (not NULL) marks no-text rows as
|
||||
// swept; NULL means pending, so COALESCE is what drains them from the
|
||||
// queue.
|
||||
_, err = sqlDB.ExecContext(ctx, `
|
||||
UPDATE chat_messages
|
||||
SET search_tsv = COALESCE(to_tsvector('simple', chat_message_search_text(content)), ''::tsvector)
|
||||
WHERE id = ANY($1)`, pq.Array([]int64{eligibleText.ID, eligibleNoText.ID}))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, pendingIDs(ctx, 10), "swept rows must leave the queue, including no-text rows")
|
||||
|
||||
// Soft-deleting an unswept row removes it from the queue without a sweep.
|
||||
unswept := newMsg(database.ChatMessageRoleUser, database.ChatMessageVisibilityBoth, textContent("unswept deploy row"))
|
||||
require.Equal(t, []int64{unswept.ID}, pendingIDs(ctx, 10))
|
||||
_, err = sqlDB.ExecContext(ctx, `UPDATE chat_messages SET deleted = true WHERE id = $1`, unswept.ID)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, pendingIDs(ctx, 10))
|
||||
|
||||
// Search contract: populate search_tsv on every row (including
|
||||
// ineligible ones) and assert the search-index predicate filters them.
|
||||
_, err = sqlDB.ExecContext(ctx, `
|
||||
UPDATE chat_messages
|
||||
SET search_tsv = COALESCE(to_tsvector('simple', chat_message_search_text(content)), ''::tsvector)
|
||||
WHERE chat_id = $1`, chat.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
rows, err := sqlDB.QueryContext(ctx, `
|
||||
SELECT id FROM chat_messages
|
||||
WHERE search_tsv @@ websearch_to_tsquery('simple', $1)
|
||||
AND search_tsv IS NOT NULL
|
||||
AND `+eligibilityPredicate+`
|
||||
ORDER BY id`, "deploy")
|
||||
require.NoError(t, err)
|
||||
defer rows.Close()
|
||||
var matched []int64
|
||||
for rows.Next() {
|
||||
var id int64
|
||||
require.NoError(t, rows.Scan(&id))
|
||||
matched = append(matched, id)
|
||||
}
|
||||
require.NoError(t, rows.Err())
|
||||
require.Equal(t, []int64{eligibleText.ID}, matched,
|
||||
"search must exclude deleted, model-only, and tool-role rows (%d %d %d)",
|
||||
toolMsg.ID, modelOnly.ID, deletedMsg.ID)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user