mirror of
https://github.com/basketikun/infinite-canvas.git
synced 2026-07-24 15:24:06 +08:00
180 lines
5.0 KiB
Go
180 lines
5.0 KiB
Go
package repository
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
|
|
"github.com/basketikun/infinite-canvas/config"
|
|
"github.com/basketikun/infinite-canvas/model"
|
|
"github.com/glebarez/sqlite"
|
|
mysqldriver "github.com/go-sql-driver/mysql"
|
|
"github.com/jackc/pgx/v5"
|
|
"github.com/jackc/pgx/v5/pgconn"
|
|
gormmysql "gorm.io/driver/mysql"
|
|
"gorm.io/driver/postgres"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
var promptCategories = []model.PromptCategory{
|
|
{Category: "system", Name: "系统", Description: "系统提示词分类"},
|
|
{Category: "gpt-image-2-prompts", Name: "GPT Image 2 Prompts", Description: "EvoLinkAI 的 GPT Image 2 案例提示词分类", GithubURL: "https://github.com/EvoLinkAI/awesome-gpt-image-2-API-and-Prompts", Remote: true},
|
|
{Category: "awesome-gpt-image", Name: "Awesome GPT Image", Description: "ZeroLu 的中文 GPT Image 提示词分类", GithubURL: "https://github.com/ZeroLu/awesome-gpt-image", Remote: true},
|
|
{Category: "awesome-gpt4o-image-prompts", Name: "Awesome GPT4o Image Prompts", Description: "ImgEdify 的 GPT-4o 图像提示词分类", GithubURL: "https://github.com/ImgEdify/Awesome-GPT4o-Image-Prompts", Remote: true},
|
|
{Category: "youmind-gpt-image-2", Name: "YouMind GPT Image 2", Description: "YouMind OpenLab 的 GPT Image 2 中文提示词分类", GithubURL: "https://github.com/YouMind-OpenLab/awesome-gpt-image-2", Remote: true},
|
|
{Category: "youmind-nano-banana-pro", Name: "YouMind Nano Banana Pro", Description: "YouMind OpenLab 的 Nano Banana Pro 中文提示词分类", GithubURL: "https://github.com/YouMind-OpenLab/awesome-nano-banana-pro-prompts", Remote: true},
|
|
{Category: "davidwu-gpt-image2-prompts", Name: "awesome-gpt-image2-prompts", Description: "davidwuw0811-boop 整理的 GPT Image 2 提示词分类", GithubURL: "https://github.com/davidwuw0811-boop/awesome-gpt-image2-prompts", Remote: true},
|
|
}
|
|
|
|
var (
|
|
db *gorm.DB
|
|
dbOnce sync.Once
|
|
dbErr error
|
|
)
|
|
|
|
// DB 初始化并返回全局数据库连接。
|
|
func DB() (*gorm.DB, error) {
|
|
dbOnce.Do(func() {
|
|
driver := strings.ToLower(strings.TrimSpace(config.Cfg.StorageDriver))
|
|
if driver == "" {
|
|
driver = "sqlite"
|
|
}
|
|
dsn := config.Cfg.DatabaseDSN
|
|
if driver == "sqlite" && dsn != ":memory:" {
|
|
_ = os.MkdirAll(filepath.Dir(dsn), 0755)
|
|
}
|
|
if isPostgresDriver(driver) {
|
|
dbErr = ensurePostgresDatabase(dsn)
|
|
if dbErr != nil {
|
|
return
|
|
}
|
|
}
|
|
if driver == "mysql" {
|
|
dbErr = ensureMySQLDatabase(dsn)
|
|
if dbErr != nil {
|
|
return
|
|
}
|
|
}
|
|
db, dbErr = gorm.Open(dialector(driver, dsn), &gorm.Config{})
|
|
if dbErr != nil {
|
|
return
|
|
}
|
|
dbErr = db.AutoMigrate(
|
|
&model.User{},
|
|
&model.CreditLog{},
|
|
&model.Prompt{},
|
|
&model.Asset{},
|
|
&model.Setting{},
|
|
)
|
|
})
|
|
return db, dbErr
|
|
}
|
|
|
|
func dialector(driver string, dsn string) gorm.Dialector {
|
|
switch driver {
|
|
case "mysql":
|
|
return gormmysql.Open(dsn)
|
|
case "postgres", "postgresql":
|
|
return postgres.Open(dsn)
|
|
default:
|
|
return sqlite.Open(dsn)
|
|
}
|
|
}
|
|
|
|
func isPostgresDriver(driver string) bool {
|
|
return driver == "postgres" || driver == "postgresql"
|
|
}
|
|
|
|
func ensureMySQLDatabase(dsn string) error {
|
|
cfg, err := mysqldriver.ParseDSN(dsn)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
target := strings.TrimSpace(cfg.DBName)
|
|
if target == "" {
|
|
return nil
|
|
}
|
|
ctx := context.Background()
|
|
targetDB, err := sql.Open("mysql", dsn)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
err = targetDB.PingContext(ctx)
|
|
_ = targetDB.Close()
|
|
if err == nil {
|
|
return nil
|
|
}
|
|
if !isMySQLError(err, 1049) {
|
|
return err
|
|
}
|
|
|
|
maintenance := cfg.Clone()
|
|
maintenance.DBName = ""
|
|
serverDB, err := sql.Open("mysql", maintenance.FormatDSN())
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer serverDB.Close()
|
|
|
|
_, err = serverDB.ExecContext(ctx, "CREATE DATABASE "+quoteMySQLIdentifier(target)+" CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci")
|
|
if isMySQLError(err, 1007) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
|
|
func ensurePostgresDatabase(dsn string) error {
|
|
cfg, err := pgx.ParseConfig(dsn)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
target := strings.TrimSpace(cfg.Database)
|
|
if target == "" {
|
|
return nil
|
|
}
|
|
ctx := context.Background()
|
|
conn, err := pgx.ConnectConfig(ctx, cfg)
|
|
if err == nil {
|
|
_ = conn.Close(ctx)
|
|
return nil
|
|
}
|
|
if !isPostgresError(err, "3D000") {
|
|
return err
|
|
}
|
|
|
|
maintenance := cfg.Copy()
|
|
maintenance.Database = "postgres"
|
|
if strings.EqualFold(target, "postgres") {
|
|
maintenance.Database = "template1"
|
|
}
|
|
conn, err = pgx.ConnectConfig(ctx, maintenance)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer conn.Close(ctx)
|
|
|
|
_, err = conn.Exec(ctx, "CREATE DATABASE "+pgx.Identifier{target}.Sanitize(), pgx.QueryExecModeExec)
|
|
if isPostgresError(err, "42P04") {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
|
|
func isMySQLError(err error, number uint16) bool {
|
|
var mysqlErr *mysqldriver.MySQLError
|
|
return errors.As(err, &mysqlErr) && mysqlErr.Number == number
|
|
}
|
|
|
|
func isPostgresError(err error, code string) bool {
|
|
var pgErr *pgconn.PgError
|
|
return errors.As(err, &pgErr) && pgErr.Code == code
|
|
}
|
|
|
|
func quoteMySQLIdentifier(name string) string {
|
|
return "`" + strings.ReplaceAll(name, "`", "``") + "`"
|
|
}
|