Files
chorus/internal/seed/seed.go
T

100 lines
3.3 KiB
Go

package seed
import (
"context"
"database/sql"
"fmt"
"net/url"
"strings"
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
)
type Options struct {
UserEmail string
UserPassword string
ProviderBaseURL string
}
func Run(ctx context.Context, db *sql.DB, options Options) error {
options.UserEmail = strings.ToLower(strings.TrimSpace(options.UserEmail))
options.ProviderBaseURL = strings.TrimSpace(options.ProviderBaseURL)
if options.UserEmail == "" || options.UserPassword == "" || options.ProviderBaseURL == "" {
return fmt.Errorf("seed options require user email, user password, and provider base URL")
}
parsedURL, err := url.Parse(options.ProviderBaseURL)
if err != nil || parsedURL.Scheme == "" || parsedURL.Host == "" {
return fmt.Errorf("seed provider base URL is invalid")
}
defaults, err := LoadDefaults()
if err != nil {
return err
}
passwordHash, err := passwordpkg.Encode(options.UserPassword)
if err != nil {
return fmt.Errorf("hash seed user password: %w", err)
}
tx, err := db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin seed transaction: %w", err)
}
defer tx.Rollback()
if _, err := tx.ExecContext(ctx, `
INSERT INTO users (email, password_hash, display_name, status)
VALUES (?, ?, ?, 'active')
ON DUPLICATE KEY UPDATE
password_hash = VALUES(password_hash), display_name = VALUES(display_name), status = 'active'`,
options.UserEmail, passwordHash, "MVP-0 Test User"); err != nil {
return fmt.Errorf("seed user: %w", err)
}
result, err := tx.ExecContext(ctx, `
INSERT INTO providers (slug, name, base_url, auth_type, api_key_enc, enabled)
VALUES (?, ?, ?, 'none', NULL, TRUE)
ON DUPLICATE KEY UPDATE
id = LAST_INSERT_ID(id), name = VALUES(name), base_url = VALUES(base_url),
auth_type = 'none', api_key_enc = NULL, enabled = TRUE`,
defaults.Provider.Slug, defaults.Provider.Name, options.ProviderBaseURL)
if err != nil {
return fmt.Errorf("seed provider: %w", err)
}
providerID, err := result.LastInsertId()
if err != nil {
return fmt.Errorf("read seeded provider id: %w", err)
}
for _, model := range defaults.Models {
if _, err := tx.ExecContext(ctx, `
INSERT INTO provider_models
(provider_id, name, model_id, api_type, kind, extra_body, timeout_ms, weight, enabled)
VALUES (?, ?, ?, ?, ?, JSON_OBJECT(), 30000, 100, TRUE)
ON DUPLICATE KEY UPDATE
name = VALUES(name), kind = VALUES(kind), extra_body = JSON_OBJECT(),
timeout_ms = VALUES(timeout_ms), weight = VALUES(weight), enabled = TRUE`,
providerID, model.Name, model.ModelID, model.APIType, model.Kind); err != nil {
return fmt.Errorf("seed provider model %s: %w", model.ModelID, err)
}
}
for _, prompt := range defaults.Prompts {
if _, err := tx.ExecContext(ctx, `
INSERT INTO prompt_templates
(template_key, kind, api_type, name, version, template_text, enabled)
VALUES (?, ?, ?, ?, ?, ?, TRUE)
ON DUPLICATE KEY UPDATE
kind = VALUES(kind), api_type = VALUES(api_type), name = VALUES(name),
version = VALUES(version), template_text = VALUES(template_text), enabled = TRUE`,
prompt.Key, prompt.Kind, prompt.APIType, prompt.Name, prompt.Version, prompt.Template); err != nil {
return fmt.Errorf("seed prompt %s: %w", prompt.Key, err)
}
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit seed transaction: %w", err)
}
return nil
}