263 lines
8.2 KiB
Go
263 lines
8.2 KiB
Go
package proxy
|
|
|
|
import (
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"log/slog"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/go-chi/chi/v5"
|
|
"github.com/go-chi/cors"
|
|
|
|
"github.com/Silo-Server/silo-server/internal/nodeconfig"
|
|
"github.com/Silo-Server/silo-server/internal/nodesessions"
|
|
"github.com/Silo-Server/silo-server/internal/playback"
|
|
"github.com/Silo-Server/silo-server/internal/streamtoken"
|
|
)
|
|
|
|
// Server is the HTTP handler for proxy mode.
|
|
type Server struct {
|
|
watcher *nodeconfig.Watcher
|
|
tracker *nodesessions.Tracker
|
|
httpClient *http.Client
|
|
}
|
|
|
|
// NewServer creates a new proxy server backed by a config watcher and session tracker.
|
|
func NewServer(watcher *nodeconfig.Watcher, tracker *nodesessions.Tracker) *Server {
|
|
return &Server{
|
|
watcher: watcher,
|
|
tracker: tracker,
|
|
httpClient: &http.Client{}, // no timeout for long streams
|
|
}
|
|
}
|
|
|
|
// Handler returns the chi.Router with all proxy routes mounted.
|
|
func (s *Server) Handler() http.Handler {
|
|
r := chi.NewRouter()
|
|
// hls.js uses XHR for manifest/segment fetches which are subject to
|
|
// CORS when the proxy runs on a different origin than the web app.
|
|
r.Use(cors.Handler(cors.Options{
|
|
AllowedOrigins: []string{"*"},
|
|
AllowedMethods: []string{"GET", "HEAD", "OPTIONS"},
|
|
AllowedHeaders: []string{"Accept", "Authorization", "Content-Type", "Range"},
|
|
MaxAge: 86400,
|
|
}))
|
|
r.Get("/api/v1/health", s.handleHealth)
|
|
r.Head("/stream/direct/{token}", s.handleDirectPlay)
|
|
r.Get("/stream/direct/{token}", s.handleDirectPlay)
|
|
r.Head("/stream/remux/{token}", s.handleRemux)
|
|
r.Get("/stream/remux/{token}", s.handleRemux)
|
|
r.Head("/stream/transcode/{token}/master.m3u8", s.handleTranscodeManifest)
|
|
r.Get("/stream/transcode/{token}/master.m3u8", s.handleTranscodeManifest)
|
|
r.Get("/stream/transcode/{token}/segment/{name}", s.handleTranscodeSegment)
|
|
r.Get("/stream/subtitles/{token}/{track}", s.handleSubtitle)
|
|
|
|
// Admin routes — bearer-auth protected.
|
|
r.Group(func(r chi.Router) {
|
|
r.Use(s.requireBearer)
|
|
r.Post("/admin/force-reload", s.handleForceReload)
|
|
r.Get("/status", s.handleStatus)
|
|
})
|
|
return r
|
|
}
|
|
|
|
type healthResponse struct {
|
|
Status string `json:"status"`
|
|
}
|
|
|
|
func (s *Server) handleHealth(w http.ResponseWriter, _ *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(healthResponse{Status: "ok"})
|
|
}
|
|
|
|
// requireBearer checks Authorization: Bearer {secret} for admin endpoints.
|
|
func (s *Server) requireBearer(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
cfg := s.watcher.Config()
|
|
auth := r.Header.Get("Authorization")
|
|
if !strings.HasPrefix(auth, "Bearer ") || strings.TrimPrefix(auth, "Bearer ") != cfg.Auth.JWTSecret {
|
|
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
|
return
|
|
}
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
// verifyToken extracts and validates the stream token from the URL.
|
|
func (s *Server) verifyToken(w http.ResponseWriter, r *http.Request) *streamtoken.Claims {
|
|
cfg := s.watcher.Config()
|
|
tokenStr := chi.URLParam(r, "token")
|
|
claims, err := streamtoken.Verify(tokenStr, cfg.Auth.JWTSecret)
|
|
if err != nil {
|
|
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
|
return nil
|
|
}
|
|
return claims
|
|
}
|
|
|
|
func (s *Server) handleDirectPlay(w http.ResponseWriter, r *http.Request) {
|
|
claims := s.verifyToken(w, r)
|
|
if claims == nil {
|
|
return
|
|
}
|
|
|
|
info := nodesessions.SessionInfo{
|
|
SessionID: claims.SessionID,
|
|
NodeURL: s.tracker.NodeURL(),
|
|
NodeName: s.tracker.NodeName(),
|
|
Type: "direct_play",
|
|
StartedAt: time.Now().UTC().Format(time.RFC3339),
|
|
}
|
|
s.tracker.Track(r.Context(), info)
|
|
defer s.tracker.Remove(r.Context(), claims.SessionID)
|
|
|
|
http.ServeFile(w, r, claims.MediaPath)
|
|
}
|
|
|
|
func (s *Server) handleRemux(w http.ResponseWriter, r *http.Request) {
|
|
claims := s.verifyToken(w, r)
|
|
if claims == nil {
|
|
return
|
|
}
|
|
|
|
info := nodesessions.SessionInfo{
|
|
SessionID: claims.SessionID,
|
|
NodeURL: s.tracker.NodeURL(),
|
|
NodeName: s.tracker.NodeName(),
|
|
Type: "remux",
|
|
StartedAt: time.Now().UTC().Format(time.RFC3339),
|
|
}
|
|
s.tracker.Track(r.Context(), info)
|
|
defer s.tracker.Remove(r.Context(), claims.SessionID)
|
|
|
|
seekSeconds := 0.0
|
|
if seekStr := r.URL.Query().Get("seek"); seekStr != "" {
|
|
if v, err := strconv.ParseFloat(seekStr, 64); err == nil {
|
|
seekSeconds = v
|
|
}
|
|
}
|
|
_ = playback.ServeRemux(w, r, claims.MediaPath, "mp4", seekSeconds, claims.TranscodeAudio, claims.AudioTrackIndex)
|
|
}
|
|
|
|
func (s *Server) handleTranscodeManifest(w http.ResponseWriter, r *http.Request) {
|
|
claims := s.verifyToken(w, r)
|
|
if claims == nil {
|
|
return
|
|
}
|
|
s.proxyToTranscodeNode(w, r, claims, "/transcode/"+claims.SessionID+"/master.m3u8")
|
|
}
|
|
|
|
func (s *Server) handleTranscodeSegment(w http.ResponseWriter, r *http.Request) {
|
|
claims := s.verifyToken(w, r)
|
|
if claims == nil {
|
|
return
|
|
}
|
|
name := chi.URLParam(r, "name")
|
|
s.proxyToTranscodeNode(w, r, claims, "/transcode/"+claims.SessionID+"/segment/"+name)
|
|
}
|
|
|
|
func (s *Server) handleSubtitle(w http.ResponseWriter, r *http.Request) {
|
|
claims := s.verifyToken(w, r)
|
|
if claims == nil {
|
|
return
|
|
}
|
|
cfg := s.watcher.Config()
|
|
trackParam := chi.URLParam(r, "track")
|
|
trackIndex, requestedFormat, err := playback.ParseSubtitleTrackParam(trackParam)
|
|
if err != nil {
|
|
http.Error(w, "invalid subtitle index", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
// When the URL requests ASS format (e.g. /subtitles/{token}/2.ass),
|
|
// extract as raw ASS to preserve styling for client-side rendering.
|
|
if requestedFormat == "ass" {
|
|
data, err := playback.ExtractSubtitleWithFormat(r.Context(), claims.MediaPath, trackIndex, "ass", cfg.Playback.FFmpegPath)
|
|
if err != nil {
|
|
slog.Error("extract subtitle (ass)", "error", err, "track", trackIndex, "path", claims.MediaPath, "playback_session_id", claims.SessionID)
|
|
http.Error(w, "subtitle extraction failed", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
playback.ServeSubtitle(w, data, "ass")
|
|
return
|
|
}
|
|
|
|
data, format, err := playback.ExtractSubtitle(r.Context(), claims.MediaPath, trackIndex, cfg.Playback.FFmpegPath)
|
|
if err != nil {
|
|
slog.Error("extract subtitle", "error", err, "track", trackIndex, "path", claims.MediaPath, "playback_session_id", claims.SessionID)
|
|
http.Error(w, "subtitle extraction failed", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
vtt, err := playback.ConvertToVTT(data, format)
|
|
if err != nil {
|
|
slog.Error("convert to vtt", "error", err, "playback_session_id", claims.SessionID)
|
|
http.Error(w, "subtitle conversion failed", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
playback.ServeSubtitle(w, vtt, "vtt")
|
|
}
|
|
|
|
// proxyToTranscodeNode forwards the request to the transcode node specified in the claims.
|
|
func (s *Server) proxyToTranscodeNode(w http.ResponseWriter, r *http.Request, claims *streamtoken.Claims, path string) {
|
|
cfg := s.watcher.Config()
|
|
if claims.TranscodeNode == "" {
|
|
http.Error(w, "no transcode node in token", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
targetURL := claims.TranscodeNode + path
|
|
if rawQuery := r.URL.RawQuery; rawQuery != "" {
|
|
targetURL = fmt.Sprintf("%s?%s", targetURL, rawQuery)
|
|
}
|
|
|
|
req, err := http.NewRequestWithContext(r.Context(), r.Method, targetURL, nil)
|
|
if err != nil {
|
|
http.Error(w, "internal error", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+cfg.Auth.JWTSecret)
|
|
|
|
resp, err := s.httpClient.Do(req)
|
|
if err != nil {
|
|
slog.Error("proxy to transcode node", "error", err, "url", targetURL, "playback_session_id", claims.SessionID)
|
|
http.Error(w, "transcode node unavailable", http.StatusBadGateway)
|
|
return
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
// Copy response headers
|
|
for k, vv := range resp.Header {
|
|
for _, v := range vv {
|
|
w.Header().Add(k, v)
|
|
}
|
|
}
|
|
w.WriteHeader(resp.StatusCode)
|
|
io.Copy(w, resp.Body)
|
|
}
|
|
|
|
func (s *Server) handleForceReload(w http.ResponseWriter, r *http.Request) {
|
|
if err := s.watcher.ForceReload(r.Context()); err != nil {
|
|
http.Error(w, "reload failed: "+err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
slog.Info("proxy force reload completed")
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}
|
|
|
|
type statusResponse struct {
|
|
ActiveSessions int `json:"active_sessions"`
|
|
}
|
|
|
|
func (s *Server) handleStatus(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(statusResponse{
|
|
ActiveSessions: s.tracker.ActiveCount(),
|
|
})
|
|
}
|