269 lines
8.9 KiB
Go
269 lines
8.9 KiB
Go
package postgres
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"encoding/json"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"git.gobha.me/xcaliber/chat-switchboard/models"
|
|
"git.gobha.me/xcaliber/chat-switchboard/store"
|
|
)
|
|
|
|
type ChannelStore struct{}
|
|
|
|
func NewChannelStore() *ChannelStore { return &ChannelStore{} }
|
|
|
|
func (s *ChannelStore) Create(ctx context.Context, ch *models.Channel) error {
|
|
return DB.QueryRowContext(ctx, `
|
|
INSERT INTO channels (user_id, title, description, type, model, system_prompt,
|
|
provider_config_id, is_archived, is_pinned, folder_id, team_id, settings)
|
|
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12)
|
|
RETURNING id, created_at, updated_at`,
|
|
ch.UserID, ch.Title, ch.Description, ch.Type, ch.Model, ch.SystemPrompt,
|
|
models.NullString(ch.ProviderConfigID), ch.IsArchived, ch.IsPinned,
|
|
models.NullString(ch.FolderID), models.NullString(ch.TeamID), ToJSON(ch.Settings),
|
|
).Scan(&ch.ID, &ch.CreatedAt, &ch.UpdatedAt)
|
|
}
|
|
|
|
func (s *ChannelStore) GetByID(ctx context.Context, id string) (*models.Channel, error) {
|
|
var ch models.Channel
|
|
var providerConfigID, folderID, teamID sql.NullString
|
|
var desc sql.NullString
|
|
var settingsJSON []byte
|
|
err := DB.QueryRowContext(ctx, `
|
|
SELECT id, user_id, title, description, type, model, system_prompt,
|
|
provider_config_id, is_archived, is_pinned, folder_id, team_id, settings,
|
|
created_at, updated_at
|
|
FROM channels WHERE id = $1`, id).Scan(
|
|
&ch.ID, &ch.UserID, &ch.Title, &desc, &ch.Type, &ch.Model, &ch.SystemPrompt,
|
|
&providerConfigID, &ch.IsArchived, &ch.IsPinned, &folderID, &teamID, &settingsJSON,
|
|
&ch.CreatedAt, &ch.UpdatedAt,
|
|
)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
ch.Description = NullableString(desc)
|
|
ch.ProviderConfigID = NullableStringPtr(providerConfigID)
|
|
ch.FolderID = NullableStringPtr(folderID)
|
|
ch.TeamID = NullableStringPtr(teamID)
|
|
json.Unmarshal(settingsJSON, &ch.Settings)
|
|
return &ch, nil
|
|
}
|
|
|
|
func (s *ChannelStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
|
b := NewUpdate("channels")
|
|
for k, v := range fields {
|
|
if k == "settings" || k == "tags" {
|
|
b.SetJSON(k, v)
|
|
} else {
|
|
b.Set(k, v)
|
|
}
|
|
}
|
|
if !b.HasSets() {
|
|
return nil
|
|
}
|
|
b.Where("id", id)
|
|
_, err := b.Exec(DB)
|
|
return err
|
|
}
|
|
|
|
func (s *ChannelStore) Delete(ctx context.Context, id string) error {
|
|
_, err := DB.ExecContext(ctx, "DELETE FROM channels WHERE id = $1", id)
|
|
return err
|
|
}
|
|
|
|
func (s *ChannelStore) ListForUser(ctx context.Context, userID string, opts store.ListOptions) ([]models.Channel, int, error) {
|
|
// Count
|
|
var total int
|
|
DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM channels WHERE user_id = $1", userID).Scan(&total)
|
|
|
|
b := NewSelect(
|
|
"id, user_id, title, description, type, model, system_prompt, provider_config_id, is_archived, is_pinned, folder_id, team_id, settings, created_at, updated_at",
|
|
"channels",
|
|
).Where("user_id = ?", userID)
|
|
|
|
if opts.Sort == "" {
|
|
b.OrderBy("updated_at", "DESC")
|
|
}
|
|
b.Paginate(opts)
|
|
|
|
q, args := b.Build()
|
|
rows, err := DB.QueryContext(ctx, q, args...)
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var result []models.Channel
|
|
for rows.Next() {
|
|
var ch models.Channel
|
|
var providerConfigID, folderID, teamID, desc sql.NullString
|
|
var settingsJSON []byte
|
|
err := rows.Scan(&ch.ID, &ch.UserID, &ch.Title, &desc, &ch.Type, &ch.Model,
|
|
&ch.SystemPrompt, &providerConfigID, &ch.IsArchived, &ch.IsPinned,
|
|
&folderID, &teamID, &settingsJSON, &ch.CreatedAt, &ch.UpdatedAt)
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
ch.Description = NullableString(desc)
|
|
ch.ProviderConfigID = NullableStringPtr(providerConfigID)
|
|
ch.FolderID = NullableStringPtr(folderID)
|
|
ch.TeamID = NullableStringPtr(teamID)
|
|
json.Unmarshal(settingsJSON, &ch.Settings)
|
|
result = append(result, ch)
|
|
}
|
|
return result, total, rows.Err()
|
|
}
|
|
|
|
func (s *ChannelStore) Search(ctx context.Context, userID, query string, opts store.ListOptions) ([]models.Channel, int, error) {
|
|
// Simple title search for now — will add full-text when search feature lands
|
|
b := NewSelect(
|
|
"id, user_id, title, description, type, model, system_prompt, provider_config_id, is_archived, is_pinned, folder_id, team_id, settings, created_at, updated_at",
|
|
"channels",
|
|
).Where("user_id = ?", userID).Where("title ILIKE ?", "%"+query+"%")
|
|
b.OrderBy("updated_at", "DESC")
|
|
b.Paginate(opts)
|
|
|
|
var total int
|
|
DB.QueryRowContext(ctx,
|
|
"SELECT COUNT(*) FROM channels WHERE user_id = $1 AND title ILIKE $2",
|
|
userID, "%"+query+"%").Scan(&total)
|
|
|
|
q, args := b.Build()
|
|
rows, err := DB.QueryContext(ctx, q, args...)
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var result []models.Channel
|
|
for rows.Next() {
|
|
var ch models.Channel
|
|
var providerConfigID, folderID, teamID, desc sql.NullString
|
|
var settingsJSON []byte
|
|
err := rows.Scan(&ch.ID, &ch.UserID, &ch.Title, &desc, &ch.Type, &ch.Model,
|
|
&ch.SystemPrompt, &providerConfigID, &ch.IsArchived, &ch.IsPinned,
|
|
&folderID, &teamID, &settingsJSON, &ch.CreatedAt, &ch.UpdatedAt)
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
ch.Description = NullableString(desc)
|
|
ch.ProviderConfigID = NullableStringPtr(providerConfigID)
|
|
ch.FolderID = NullableStringPtr(folderID)
|
|
ch.TeamID = NullableStringPtr(teamID)
|
|
json.Unmarshal(settingsJSON, &ch.Settings)
|
|
result = append(result, ch)
|
|
}
|
|
return result, total, rows.Err()
|
|
}
|
|
|
|
func (s *ChannelStore) GetCursor(ctx context.Context, channelID, userID string) (*models.ChannelCursor, error) {
|
|
var c models.ChannelCursor
|
|
var leafID sql.NullString
|
|
err := DB.QueryRowContext(ctx, `
|
|
SELECT id, channel_id, user_id, active_leaf_id, updated_at
|
|
FROM channel_cursors WHERE channel_id = $1 AND user_id = $2`,
|
|
channelID, userID).Scan(&c.ID, &c.ChannelID, &c.UserID, &leafID, &c.UpdatedAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
c.ActiveLeafID = NullableStringPtr(leafID)
|
|
return &c, nil
|
|
}
|
|
|
|
func (s *ChannelStore) SetCursor(ctx context.Context, channelID, userID, leafID string) error {
|
|
_, err := DB.ExecContext(ctx, `
|
|
INSERT INTO channel_cursors (channel_id, user_id, active_leaf_id)
|
|
VALUES ($1, $2, $3)
|
|
ON CONFLICT (channel_id, user_id) DO UPDATE SET active_leaf_id = $3, updated_at = NOW()`,
|
|
channelID, userID, leafID)
|
|
return err
|
|
}
|
|
|
|
func (s *ChannelStore) SetModel(ctx context.Context, cm *models.ChannelModel) error {
|
|
_, err := DB.ExecContext(ctx, `
|
|
INSERT INTO channel_models (channel_id, model_id, provider_config_id, display_name, system_prompt, settings, is_default)
|
|
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
|
ON CONFLICT (channel_id, model_id) DO UPDATE SET
|
|
provider_config_id = $3, display_name = $4, system_prompt = $5, settings = $6, is_default = $7`,
|
|
cm.ChannelID, cm.ModelID, cm.ProviderConfigID, cm.DisplayName, cm.SystemPrompt, "{}", cm.IsDefault)
|
|
return err
|
|
}
|
|
|
|
func (s *ChannelStore) GetModels(ctx context.Context, channelID string) ([]models.ChannelModel, error) {
|
|
rows, err := DB.QueryContext(ctx, `
|
|
SELECT id, channel_id, model_id, COALESCE(provider_config_id::text, ''),
|
|
COALESCE(display_name, ''), COALESCE(system_prompt, ''), is_default
|
|
FROM channel_models WHERE channel_id = $1`, channelID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var result []models.ChannelModel
|
|
for rows.Next() {
|
|
var cm models.ChannelModel
|
|
if err := rows.Scan(&cm.ID, &cm.ChannelID, &cm.ModelID, &cm.ProviderConfigID,
|
|
&cm.DisplayName, &cm.SystemPrompt, &cm.IsDefault); err != nil {
|
|
return nil, err
|
|
}
|
|
result = append(result, cm)
|
|
}
|
|
return result, rows.Err()
|
|
}
|
|
|
|
func (s *ChannelStore) GetModelByID(ctx context.Context, id string) (*models.ChannelModel, error) {
|
|
var cm models.ChannelModel
|
|
err := DB.QueryRowContext(ctx, `
|
|
SELECT id, channel_id, model_id, COALESCE(provider_config_id::text, ''),
|
|
COALESCE(display_name, ''), COALESCE(system_prompt, ''), is_default
|
|
FROM channel_models WHERE id = $1`, id).Scan(
|
|
&cm.ID, &cm.ChannelID, &cm.ModelID, &cm.ProviderConfigID,
|
|
&cm.DisplayName, &cm.SystemPrompt, &cm.IsDefault)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &cm, nil
|
|
}
|
|
|
|
func (s *ChannelStore) UpdateModel(ctx context.Context, id string, fields map[string]interface{}) error {
|
|
if len(fields) == 0 {
|
|
return nil
|
|
}
|
|
// Build SET clause dynamically
|
|
sets := make([]string, 0, len(fields))
|
|
args := make([]interface{}, 0, len(fields)+1)
|
|
i := 1
|
|
for col, val := range fields {
|
|
sets = append(sets, col+" = $"+fmt.Sprintf("%d", i))
|
|
args = append(args, val)
|
|
i++
|
|
}
|
|
args = append(args, id)
|
|
query := "UPDATE channel_models SET " + strings.Join(sets, ", ") + " WHERE id = $" + fmt.Sprintf("%d", i)
|
|
_, err := DB.ExecContext(ctx, query, args...)
|
|
return err
|
|
}
|
|
|
|
func (s *ChannelStore) DeleteModel(ctx context.Context, id string) error {
|
|
res, err := DB.ExecContext(ctx, `DELETE FROM channel_models WHERE id = $1`, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
n, _ := res.RowsAffected()
|
|
if n == 0 {
|
|
return sql.ErrNoRows
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *ChannelStore) UserOwns(ctx context.Context, channelID, userID string) (bool, error) {
|
|
var exists bool
|
|
err := DB.QueryRowContext(ctx,
|
|
"SELECT EXISTS(SELECT 1 FROM channels WHERE id = $1 AND user_id = $2)",
|
|
channelID, userID).Scan(&exists)
|
|
return exists, err
|
|
}
|