175 lines
5.0 KiB
Go
175 lines
5.0 KiB
Go
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()
|
||
}
|