streaming

Media streaming and broadcast systems in Go
Log | Files | Refs | README | LICENSE

hlsserve.go (4299B)


      1 // Command hlsserve serves a live HLS stream from MPEG-TS video received over TCP.
      2 // The options are:
      3 //
      4 //	-l address
      5 //		Listen for MPEG-TS streams on address, in host:port format. The
      6 //		default is ":9000".
      7 //	-h address
      8 //		Listen for HTTP clients on address, in host:port format. The
      9 //		default is ":8080".
     10 package main
     11 
     12 import (
     13 	"bytes"
     14 	"flag"
     15 	"fmt"
     16 	"io"
     17 	"log"
     18 	"net"
     19 	"net/http"
     20 	"os"
     21 	"path"
     22 	"path/filepath"
     23 	"strings"
     24 	"time"
     25 
     26 	"github.com/untangledco/streaming/m3u8"
     27 	"github.com/untangledco/streaming/mpegts"
     28 )
     29 
     30 func init() {
     31 	log.SetFlags(0)
     32 	log.SetPrefix("hlsserve: ")
     33 	var err error
     34 	cacheDir, err = os.UserCacheDir()
     35 	if err != nil {
     36 		log.Fatalln("user cache dir:", err)
     37 	}
     38 	cacheDir = filepath.Join(cacheDir, "hlsserve")
     39 }
     40 
     41 // rule of thumb for UDP transport
     42 const maxTSBytes int = 7 * mpegts.PacketSize
     43 
     44 const segmentDuration = 3 * time.Second
     45 
     46 var cacheDir string
     47 var sequence int
     48 
     49 func removeOld(dir string, maxAge time.Duration) error {
     50 	ents, err := os.ReadDir(dir)
     51 	if err != nil {
     52 		return err
     53 	}
     54 	for _, dent := range ents {
     55 		if !strings.HasSuffix(dent.Name(), ".ts") {
     56 			continue
     57 		}
     58 		stat, err := dent.Info()
     59 		if err != nil {
     60 			return err
     61 		}
     62 		if time.Since(stat.ModTime()) > maxAge {
     63 			if err := os.Remove(filepath.Join(dir, dent.Name())); err != nil {
     64 				return err
     65 			}
     66 			sequence++
     67 		}
     68 	}
     69 	return nil
     70 }
     71 
     72 func writeSegments(dir string, r io.Reader, ch <-chan time.Time) error {
     73 	var segment int
     74 	segments := &bytes.Buffer{}
     75 	sc := mpegts.NewScanner(r)
     76 	for {
     77 		select {
     78 		case <-ch:
     79 			s := fmt.Sprintf("%04d.ts", segment)
     80 			fname := path.Join(dir, s)
     81 			if err := os.WriteFile(fname, segments.Bytes(), 0644); err != nil {
     82 				return err
     83 			}
     84 			segments.Reset()
     85 			segment++
     86 		default:
     87 			if sc.Scan() {
     88 				if err := mpegts.Encode(segments, sc.Packet()); err != nil {
     89 					return fmt.Errorf("segment %d: encode packet: %w", segment, err)
     90 				}
     91 			}
     92 			if sc.Err() != nil {
     93 				return fmt.Errorf("segment %d: scan: %w", segment, sc.Err())
     94 			}
     95 		}
     96 	}
     97 	fmt.Fprintln(os.Stderr, "shouldn't really be returning here...?")
     98 	return nil
     99 }
    100 
    101 func makePlaylist(dir string) (*m3u8.Playlist, error) {
    102 	names, err := filepath.Glob(filepath.Join(dir, "*.ts"))
    103 	if err != nil {
    104 		return nil, fmt.Errorf("find segments: %w", err)
    105 	}
    106 	playlist := &m3u8.Playlist{
    107 		Version:        7,
    108 		TargetDuration: segmentDuration,
    109 		Sequence:       sequence,
    110 	}
    111 	for _, name := range names {
    112 		seg := m3u8.Segment{URI: path.Base(name), Duration: playlist.TargetDuration}
    113 		playlist.Segments = append(playlist.Segments, seg)
    114 	}
    115 	return playlist, nil
    116 }
    117 
    118 const usage string = "usage: hlsserve dir"
    119 
    120 func servePlaylist(dir string) http.HandlerFunc {
    121 	return func(w http.ResponseWriter, req *http.Request) {
    122 		playlist, err := makePlaylist(dir)
    123 		if err != nil {
    124 			http.Error(w, err.Error(), http.StatusInternalServerError)
    125 			return
    126 		}
    127 		w.Header().Set("Content-Type", m3u8.MimeType)
    128 		if err := m3u8.Encode(w, playlist); err != nil {
    129 			log.Printf("encode playlist: %v", err)
    130 		}
    131 	}
    132 }
    133 
    134 func setCache(seconds int, next http.Handler) http.HandlerFunc {
    135 	return func(w http.ResponseWriter, req *http.Request) {
    136 		w.Header().Set("Cache-Control", fmt.Sprintf("max-age=%d", seconds))
    137 		next.ServeHTTP(w, req)
    138 	}
    139 }
    140 
    141 var listen = ":9000"
    142 
    143 func init() {
    144 	flag.StringVar(&listen, "l", ":9000", "listen")
    145 	flag.Parse()
    146 }
    147 
    148 func main() {
    149 	if len(flag.Args()) > 1 {
    150 		fmt.Fprintln(os.Stderr, usage)
    151 		os.Exit(2)
    152 	} else if len(flag.Args()) == 1 {
    153 		cacheDir = flag.Args()[0]
    154 	}
    155 	if err := os.MkdirAll(cacheDir, 0755); err != nil {
    156 		log.Fatal(err)
    157 	}
    158 
    159 	ln, err := net.Listen("tcp", listen)
    160 	if err != nil {
    161 		log.Fatal(err)
    162 	}
    163 	conn, err := ln.Accept()
    164 	if err != nil {
    165 		log.Fatal(err)
    166 	}
    167 
    168 	go func() {
    169 		ticker := time.NewTicker(segmentDuration)
    170 		if err := writeSegments(cacheDir, conn, ticker.C); err != nil {
    171 			log.Fatalln("write segments:", err)
    172 		}
    173 	}()
    174 
    175 	go func() {
    176 		ticker := time.NewTicker(segmentDuration)
    177 		for {
    178 			select {
    179 			case <-ticker.C:
    180 				if err := removeOld(cacheDir, 3*8*segmentDuration); err != nil {
    181 					log.Println("remove old segments:", err)
    182 				}
    183 			}
    184 		}
    185 	}()
    186 
    187 	http.Handle("/playlist.m3u8", servePlaylist(cacheDir))
    188 	fsys := http.FileServer(http.FS(os.DirFS(cacheDir)))
    189 	http.Handle("/", setCache(60, fsys))
    190 	log.Fatal(http.ListenAndServe(":8000", nil))
    191 }