package database import ( "context" "database/sql" "encoding/json" "errors" "fmt" "os" "path/filepath" "strings" "entgo.io/ent/dialect" entsql "entgo.io/ent/dialect/sql" "home-vue-go/internal/config" "home-vue-go/internal/ent" "home-vue-go/internal/ent/migrate" _ "github.com/mattn/go-sqlite3" "golang.org/x/crypto/bcrypt" ) const siteConfigPayloadColumn = "config_json" type Database struct { Client *ent.Client SQL *sql.DB } func Init(dbPath string, cfg *config.Config) (*Database, error) { db, err := sql.Open("sqlite3", dbPath+"?_fk=1") if err != nil { return nil, err } drv := entsql.OpenDB(dialect.SQLite, db) client := ent.NewClient(ent.Driver(drv)) ctx := context.Background() if err := client.Schema.Create(ctx, migrate.WithForeignKeys(false)); err != nil { _ = db.Close() return nil, fmt.Errorf("数据库迁移失败: %w", err) } if err := ensureSiteConfigPayload(db); err != nil { _ = client.Close() _ = db.Close() return nil, err } database := &Database{Client: client, SQL: db} if err := initDefaultData(ctx, client); err != nil { _ = database.Close() return nil, fmt.Errorf("初始化默认数据失败: %w", err) } if err := initializeSiteSettings(ctx, database, cfg); err != nil { _ = database.Close() return nil, fmt.Errorf("初始化站点配置载荷失败: %w", err) } return database, nil } func (d *Database) Close() error { if d.Client != nil { return d.Client.Close() } if d.SQL != nil { return d.SQL.Close() } return nil } func ensureSiteConfigPayload(db *sql.DB) error { rows, err := db.Query("PRAGMA table_info(site_configs)") if err != nil { return err } defer rows.Close() for rows.Next() { var cid int var name, columnType string var notNull, primaryKey int var defaultValue any if err := rows.Scan(&cid, &name, &columnType, ¬Null, &defaultValue, &primaryKey); err != nil { return err } if name == siteConfigPayloadColumn { return rows.Err() } } if err := rows.Err(); err != nil { return err } _, err = db.Exec(`ALTER TABLE site_configs ADD COLUMN config_json TEXT NOT NULL DEFAULT '{}'`) return err } func initDefaultData(ctx context.Context, client *ent.Client) error { if _, err := client.SiteConfig.Get(ctx, 1); err != nil { if !ent.IsNotFound(err) { return err } if _, err := client.SiteConfig.Create(). SetSiteName("个人主页"). SetSiteURL("https://example.com"). SetSiteIcon("/favicon.ico"). SetSiteDescription("一个基于Vue3的个人主页"). SetSiteKeywords("个人主页,Vue3"). SetUserName("用户"). SetProfileImageURL(""). SetIcpNumber(""). SetPoliceNumber(""). SetPageTitle("个人主页"). SetFavicon("/favicon.ico"). SetUmamiScript(""). SetUmamiWebsiteID(""). SetIconLibrary(""). SetFontLibrary(""). Save(ctx); err != nil { return err } } userCount, err := client.User.Query().Count(ctx) if err != nil { return err } if userCount == 0 { hashedPassword, err := bcrypt.GenerateFromPassword([]byte("admin123"), bcrypt.DefaultCost) if err != nil { return err } _, err = client.User.Create(). SetUsername("admin"). SetPassword(string(hashedPassword)). Save(ctx) return err } return nil } func initializeSiteSettings(ctx context.Context, db *Database, cfg *config.Config) error { payload, err := db.loadSiteConfigPayload(ctx) if err != nil { return err } if strings.TrimSpace(payload) != "" && strings.TrimSpace(payload) != "{}" { var settings config.SiteSettings if err := json.Unmarshal([]byte(payload), &settings); err != nil { return fmt.Errorf("解析站点配置载荷失败: %w", err) } return nil } siteCfg, err := db.Client.SiteConfig.Get(ctx, 1) if err != nil { return err } settings := config.DefaultSiteSettings() settings.SiteName = siteCfg.SiteName settings.SiteURL = siteCfg.SiteURL settings.SiteIcon = siteCfg.SiteIcon settings.SiteDescription = siteCfg.SiteDescription settings.SiteKeywords = siteCfg.SiteKeywords settings.UserName = siteCfg.UserName settings.ProfileImageURL = siteCfg.ProfileImageURL settings.ICPNumber = siteCfg.IcpNumber settings.PoliceNumber = siteCfg.PoliceNumber settings.PageTitle = siteCfg.PageTitle settings.Favicon = siteCfg.Favicon settings.IconLibrary = siteCfg.IconLibrary settings.FontLibrary = siteCfg.FontLibrary settings.UmamiScript = siteCfg.UmamiScript settings.UmamiWebsiteID = siteCfg.UmamiWebsiteID if cfg != nil { if err := importLegacyJSON(cfg.DataDir, settings); err != nil { return err } } settings.Normalize() return db.SaveSiteSettings(ctx, settings) } func importLegacyJSON(dataDir string, settings *config.SiteSettings) error { if dataDir == "" { return nil } var footer struct { Start string `json:"start"` End string `json:"end"` } if err := readOptionalJSON(filepath.Join(dataDir, "footer_year.json"), &footer); err != nil { return err } if footer.Start != "" || footer.End != "" { settings.FooterYearStart, settings.FooterYearEnd = footer.Start, footer.End } var timer config.VisitTimerConfig if err := readOptionalJSON(filepath.Join(dataDir, "visit_timer.json"), &timer); err != nil { return err } if _, err := os.Stat(filepath.Join(dataDir, "visit_timer.json")); err == nil { settings.ShowVisitTimer = timer.ShowVisitTimer } var texts config.RotatingTextsConfig if err := readOptionalJSON(filepath.Join(dataDir, "rotating_texts.json"), &texts); err != nil { return err } if len(texts.Texts) > 0 { settings.RotatingTexts = texts.Texts } return nil } func readOptionalJSON(path string, target any) error { data, err := os.ReadFile(path) if os.IsNotExist(err) { return nil } if err != nil { return err } if err := json.Unmarshal(data, target); err != nil { return fmt.Errorf("解析 %s 失败: %w", filepath.Base(path), err) } return nil } func (d *Database) loadSiteConfigPayload(ctx context.Context) (string, error) { var payload string err := d.SQL.QueryRowContext(ctx, "SELECT config_json FROM site_configs WHERE id = 1").Scan(&payload) if errors.Is(err, sql.ErrNoRows) { return "", nil } return payload, err } func (d *Database) LoadSiteSettings(ctx context.Context) (*config.SiteSettings, error) { payload, err := d.loadSiteConfigPayload(ctx) if err != nil { return nil, err } settings := config.DefaultSiteSettings() if strings.TrimSpace(payload) != "" && strings.TrimSpace(payload) != "{}" { if err := json.Unmarshal([]byte(payload), settings); err != nil { return nil, err } } settings.Normalize() return settings, nil } func (d *Database) SaveSiteSettings(ctx context.Context, settings *config.SiteSettings) error { if settings == nil { return errors.New("site settings cannot be nil") } settings.Normalize() payload, err := json.Marshal(settings) if err != nil { return err } result, err := d.SQL.ExecContext(ctx, "UPDATE site_configs SET config_json = ? WHERE id = 1", string(payload)) if err != nil { return err } if affected, err := result.RowsAffected(); err != nil { return err } else if affected != 1 { return errors.New("site config row does not exist") } return nil }