Files
chorus/internal/seed/seed.go
T

123 lines
4.5 KiB
Go

package seed
import (
"context"
"database/sql"
"fmt"
"net/url"
"regexp"
"strings"
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
)
type Options struct {
UserUsername string
UserEmail string
UserPassword string
ProviderBaseURL string
}
func Run(ctx context.Context, db *sql.DB, options Options) error {
options.UserUsername = strings.ToLower(strings.TrimSpace(options.UserUsername))
options.UserEmail = strings.ToLower(strings.TrimSpace(options.UserEmail))
options.ProviderBaseURL = strings.TrimSpace(options.ProviderBaseURL)
if !validUsername(options.UserUsername) {
return fmt.Errorf("seed user username must be 3-64 lowercase letters, digits, dots, underscores, or hyphens and start with a letter or digit")
}
if options.UserEmail == "" || options.UserPassword == "" || options.ProviderBaseURL == "" {
return fmt.Errorf("seed options require user username, 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 (username, email, password_hash, display_name, status)
VALUES (?, ?, ?, ?, 'active')
ON DUPLICATE KEY UPDATE
username = VALUES(username), email = VALUES(email), password_hash = VALUES(password_hash),
display_name = VALUES(display_name), status = 'active'`,
options.UserUsername, 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 {
result, 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
id = LAST_INSERT_ID(id), 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)
if err != nil {
return fmt.Errorf("seed provider model %s: %w", model.ModelID, err)
}
providerModelID, err := result.LastInsertId()
if err != nil {
return fmt.Errorf("read seeded provider model %s: %w", model.ModelID, err)
}
if _, err := tx.ExecContext(ctx, `
INSERT INTO provider_model_capabilities (provider_model_id, capability)
VALUES (?, ?)
ON DUPLICATE KEY UPDATE capability = VALUES(capability)`, providerModelID, model.Capability); err != nil {
return fmt.Errorf("seed provider model capability %s: %w", model.ModelID, err)
}
}
for _, prompt := range defaults.Prompts {
if _, err := tx.ExecContext(ctx, `
INSERT INTO prompt_templates
(template_key, kind, api_type, capability, name, version, template_text, default_role_rule, enabled)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, TRUE)
ON DUPLICATE KEY UPDATE
kind = VALUES(kind), api_type = VALUES(api_type), capability = VALUES(capability), name = VALUES(name),
version = VALUES(version), template_text = VALUES(template_text),
default_role_rule = VALUES(default_role_rule), enabled = TRUE`,
prompt.Key, prompt.Kind, prompt.APIType, prompt.Capability, prompt.Name, prompt.Version, prompt.Template, prompt.DefaultRoleRule); 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
}
var usernamePattern = regexp.MustCompile(`^[a-z0-9][a-z0-9._-]{2,63}$`)
func validUsername(value string) bool { return usernamePattern.MatchString(value) }