live-sfu-demo/internal/db/db.go

175 lines
5.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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()
}