100 lines
3.3 KiB
Go
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
|
|
}
|