package db import ( "context" "database/sql" "fmt" "log" "os" "path/filepath" "strings" "time" _ "github.com/tursodatabase/go-libsql" "gospeak-live-sfu-demo/gen" ) // Open 打开 Turso 嵌入式数据库(libSQL file)。 // 仅支持嵌入模式:file:./data/live-sfu.db / :memory: / file::memory: // 已禁用远程 libsql://,如需远程请另行扩展。 func Open(dsn string) (*sql.DB, error) { if strings.HasPrefix(dsn, "libsql://") || strings.HasPrefix(dsn, "https://") { return nil, fmt.Errorf("embedded mode only: remote Turso DSN not supported (%q), use file: DSN", dsn) } if dsn == "" { dsn = "file:./data/live-sfu.db?cache=shared&_journal_mode=WAL" } // 确保本地文件目录存在 if strings.HasPrefix(dsn, "file:") { pathPart := strings.TrimPrefix(dsn, "file:") if idx := strings.Index(pathPart, "?"); idx >= 0 { pathPart = pathPart[:idx] } if pathPart != "" && pathPart != ":memory:" && pathPart != ":memory" { dir := filepath.Dir(pathPart) if dir != "." && dir != "" && dir != "/" { if err := os.MkdirAll(dir, 0755); err != nil { log.Printf("[db] mkdir %s: %v", dir, err) } } } } db, err := sql.Open("libsql", dsn) if err != nil { return nil, fmt.Errorf("open db %q: %w", dsn, err) } db.SetMaxOpenConns(1) db.SetMaxIdleConns(1) db.SetConnMaxLifetime(time.Hour) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if err := db.PingContext(ctx); err != nil { return nil, fmt.Errorf("ping db %q: %w", dsn, err) } if err := Migrate(db); err != nil { return nil, fmt.Errorf("migrate: %w", err) } log.Printf("[db] turso embedded opened: %s", dsn) return db, nil } // Migrate 创建所需表结构(幂等)。 func Migrate(db *sql.DB) error { stmts := []string{ `CREATE TABLE IF NOT EXISTS stream_targets ( room TEXT NOT NULL, backend TEXT NOT NULL, session_id TEXT, stream TEXT, publish_token TEXT, url TEXT, published_at INTEGER, PRIMARY KEY (room, backend) )`, `CREATE INDEX IF NOT EXISTS idx_stream_targets_room ON stream_targets(room)`, `CREATE TABLE IF NOT EXISTS rooms ( name TEXT PRIMARY KEY, created_at INTEGER, updated_at INTEGER )`, } for _, s := range stmts { if _, err := db.Exec(s); err != nil { return fmt.Errorf("exec %q: %w", s, err) } } return nil } func StringToBackendKind(s string) gen.BackendKind { switch strings.ToLower(strings.TrimSpace(s)) { case "cloudflare": return gen.BackendKind_BACKEND_KIND_CLOUDFLARE case "srs": return gen.BackendKind_BACKEND_KIND_SRS default: return gen.BackendKind_BACKEND_KIND_UNSPECIFIED } } // SaveTarget 持久化单个分发目标(INSERT OR REPLACE)。 func SaveTarget(ctx context.Context, db *sql.DB, room, backend string, t *gen.StreamTarget) error { if db == nil { return nil } _, err := db.ExecContext(ctx, `INSERT OR REPLACE INTO stream_targets (room, backend, session_id, stream, publish_token, url, published_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, room, backend, t.GetSessionId(), t.GetStream(), t.GetPublishToken(), t.GetUrl(), t.GetPublishedAt()) if err != nil { return err } _, _ = db.ExecContext(ctx, `INSERT OR IGNORE INTO rooms (name, created_at, updated_at) VALUES (?, ?, ?)`, room, t.GetPublishedAt(), t.GetPublishedAt()) _, _ = db.ExecContext(ctx, `UPDATE rooms SET updated_at=? WHERE name=?`, t.GetPublishedAt(), room) return nil } // DeleteTarget 删除指定房间在某后端的目标。 func DeleteTarget(ctx context.Context, db *sql.DB, room, backend string) error { if db == nil { return nil } _, err := db.ExecContext(ctx, `DELETE FROM stream_targets WHERE room=? AND backend=?`, room, backend) if err != nil { return err } var cnt int if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM stream_targets WHERE room=?`, room).Scan(&cnt); err == nil && cnt == 0 { _, _ = db.ExecContext(ctx, `DELETE FROM rooms WHERE name=?`, room) } return nil } // LoadAll 读取所有房间的全部目标,按 room 分组。 func LoadAll(ctx context.Context, db *sql.DB) (map[string]map[string]*gen.StreamTarget, error) { if db == nil { return map[string]map[string]*gen.StreamTarget{}, nil } rows, err := db.QueryContext(ctx, `SELECT room, backend, session_id, stream, publish_token, url, published_at FROM stream_targets`) if err != nil { return nil, err } defer rows.Close() out := map[string]map[string]*gen.StreamTarget{} for rows.Next() { var room, backend string var sp, se, pu, ur sql.NullString var pa sql.NullInt64 if err := rows.Scan(&room, &backend, &sp, &se, &pu, &ur, &pa); err != nil { return nil, err } target := &gen.StreamTarget{ Backend: StringToBackendKind(backend), } if sp.Valid { target.SessionId = sp.String } if se.Valid { target.Stream = se.String } if pu.Valid { target.PublishToken = pu.String } if ur.Valid { target.Url = ur.String } if pa.Valid { target.PublishedAt = pa.Int64 } if _, ok := out[room]; !ok { out[room] = map[string]*gen.StreamTarget{} } out[room][backend] = target } return out, rows.Err() }