11 Commits
17 changed files with 1610 additions and 249 deletions
+14 -7
View File
@@ -1,12 +1,7 @@
HOST=0.0.0.0 HOST=0.0.0.0
PORT=22 PORT=22
# Generation settings # Playback settings
# Stored frame resolution in pixels. Higher = sharper but bigger .tsf files.
# regenerate your frames after changing it.
FRAME_RESOLUTION=512
# Playback settings.
# Playthroughs before the session is closed (0, unlimited). # Playthroughs before the session is closed (0, unlimited).
MAX_LOOP=5 MAX_LOOP=5
# Whether to keep looping the same frame set or pick a random one after each # Whether to keep looping the same frame set or pick a random one after each
@@ -30,7 +25,17 @@ INVERT=false
FORCE_GRAYSCALE=false FORCE_GRAYSCALE=false
# Max rendered width/height in characters. # Max rendered width/height in characters.
MAX_DIMENSION=1080 MAX_DIMENSION=512
# Max rendered area (columns x rows); larger terminals are scaled down.
MAX_TERMINAL_CELLS=256000
# Memory budget in MB for the rendered-frame cache (0 disables caching).
RENDER_CACHE_MB=256
# Go soft memory limit; set below your container limit to avoid OOM kills.
GOMEMLIMIT=1GiB
# Docker Compose container memory limit. Leave headroom for mapped frame files.
MEMORY_LIMIT=2g
# Connection limits. New connections over a limit are dropped immediately. # Connection limits. New connections over a limit are dropped immediately.
# Max simultaneous connections from a single client IP. # Max simultaneous connections from a single client IP.
@@ -42,6 +47,8 @@ MAX_TOTAL_CONNECTIONS=1000
MAX_AUTH_ATTEMPTS=3 MAX_AUTH_ATTEMPTS=3
# SSH handshake deadline in ms (0 to disable). # SSH handshake deadline in ms (0 to disable).
HANDSHAKE_TIMEOUT=30000 HANDSHAKE_TIMEOUT=30000
# Maximum session lifetime in ms (0 to disable).
SESSION_TIMEOUT=600000
# Log attempted usernames/passwords. # Log attempted usernames/passwords.
LOG_CREDENTIALS=true LOG_CREDENTIALS=true
+12 -5
View File
@@ -25,7 +25,7 @@ Generate a frame set from a video through container image
```sh ```sh
docker run --rm -v ./video.mp4:/home/app/video.mp4 -v ./frames:/home/app/frames \ docker run --rm -v ./video.mp4:/home/app/video.mp4 -v ./frames:/home/app/frames \
ghcr.io/yuzuzensai/trollssh:latest trollssh --generate --video video.mp4 ghcr.io/yuzuzensai/trollssh:v1.0.1 trollssh --generate --video video.mp4 --resolution 512
``` ```
This writes `frames/<name>.tsf`, a simple container of color JPEG frames plus This writes `frames/<name>.tsf`, a simple container of color JPEG frames plus
@@ -51,13 +51,20 @@ ssh anyone@localhost
## Configuration ## Configuration
Configuration is via environment variables, loaded from a `.env` file if one Server configuration is via environment variables, loaded from a `.env` file
exists (see [`.env.example`](.env.example) for the full annotated list). if one exists (see [`.env.example`](.env.example) for the full annotated
Durations are in milliseconds. list). Durations are in milliseconds.
Host keys (`data/id_rsa`, `data/id_ed25519`) are generated on first run and Host keys (`data/id_rsa`, `data/id_ed25519`) are generated on first run and
reused afterwards. reused afterwards.
Frame generation is configured with flags:
| Flag | Default | Description |
| -------------------- | ------- | -------------------------------------------------- |
| `--generate`, `-g` | | Generate a `.tsf` frame set instead of serving |
| `--video`, `-v` | | Source video path |
| `--resolution`, `-r` | `512` | Stored frame max dimension in pixels. Higher = sharper but bigger `.tsf` files and slower rendering |
## Customization ## Customization
@@ -75,7 +82,7 @@ Requirements: Go 1.25+ and `ffmpeg` / `ffprobe` on `PATH` (only for
`--generate`). `--generate`).
```sh ```sh
go run ./src --generate --video video.mp4 go run ./src --generate --video video.mp4 --resolution 512
go run ./src go run ./src
``` ```
+2 -1
View File
@@ -1,8 +1,9 @@
services: services:
trollssh: trollssh:
image: ghcr.io/yuzuzensai/trollssh:latest image: ghcr.io/yuzuzensai/trollssh:v1.0.1
container_name: trollssh container_name: trollssh
restart: unless-stopped restart: unless-stopped
mem_limit: ${MEMORY_LIMIT:-2g}
ports: ports:
- "22:22" - "22:22"
env_file: env_file:
+119 -23
View File
@@ -1,9 +1,13 @@
package main package main
import ( import (
"bytes"
"image"
"image/jpeg"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"testing" "testing"
"time" "time"
) )
@@ -56,8 +60,9 @@ func TestTSFInvalid(t *testing.T) {
} }
// Valid container but fps <= 0. // Valid container but fps <= 0.
if err := writeTSF(path, &FramesContainer{ColorFrames: [][]byte{{1}}, FPS: 0}); err != nil { rawInvalidFPS := append(tsfHeader(0, 1), 1, 0, 0, 0, 1)
t.Fatalf("writeTSF: %v", err) if err := os.WriteFile(path, rawInvalidFPS, 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
} }
if _, err := loadTSF(path); err == nil { if _, err := loadTSF(path); err == nil {
t.Error("expected error for fps<=0") t.Error("expected error for fps<=0")
@@ -168,7 +173,11 @@ func TestFrameToAscii(t *testing.T) {
// Below threshold -> first ramp char; full brightness -> last. // Below threshold -> first ramp char; full brightness -> last.
opts := asciiOptions{brightnessThreshold: 40, charset: "standard"} opts := asciiOptions{brightnessThreshold: 40, charset: "standard"}
ramp := []rune(resolveCharset("standard")) ramp := []rune(resolveCharset("standard"))
out := []rune(frameToAscii([]byte{0, 255}, opts)) img := &image.RGBA{
Pix: []byte{0, 0, 0, 255, 255, 255, 255, 255},
Stride: 8, Rect: image.Rect(0, 0, 2, 1),
}
out := []rune(string(frameToAscii(img, buildRampLUT(ramp, opts))))
if out[0] != ramp[0] { if out[0] != ramp[0] {
t.Errorf("dark px = %q, want %q", out[0], ramp[0]) t.Errorf("dark px = %q, want %q", out[0], ramp[0])
} }
@@ -180,38 +189,125 @@ func TestFrameToAscii(t *testing.T) {
func TestFrameToAsciiInvert(t *testing.T) { func TestFrameToAsciiInvert(t *testing.T) {
opts := asciiOptions{brightnessThreshold: 40, charset: "standard", invert: true} opts := asciiOptions{brightnessThreshold: 40, charset: "standard", invert: true}
ramp := []rune(resolveCharset("standard")) ramp := []rune(resolveCharset("standard"))
out := []rune(frameToAscii([]byte{255}, opts)) img := &image.RGBA{
Pix: []byte{255, 255, 255, 255},
Stride: 4, Rect: image.Rect(0, 0, 1, 1),
}
out := []rune(string(frameToAscii(img, buildRampLUT(ramp, opts))))
if out[0] != ramp[0] { if out[0] != ramp[0] {
t.Errorf("inverted bright = %q, want %q", out[0], ramp[0]) t.Errorf("inverted bright = %q, want %q", out[0], ramp[0])
} }
} }
func TestConnectionTracker(t *testing.T) { func TestRenderConcurrentSameKey(t *testing.T) {
tr := newConnectionTracker() var jpegBuf bytes.Buffer
tr.increment("1.2.3.4") src := image.NewRGBA(image.Rect(0, 0, 16, 16))
tr.increment("1.2.3.4") for i := range src.Pix {
if !tr.hasReachedLimits("1.2.3.4", 2, 100) { src.Pix[i] = byte(i * 7)
t.Error("expected per-ip limit reached")
} }
tr.decrement("1.2.3.4") if err := jpeg.Encode(&jpegBuf, src, nil); err != nil {
tr.decrement("1.2.3.4") t.Fatal(err)
if tr.totalCount() != 0 { }
t.Errorf("total = %d", tr.totalCount())
r := newFrameRenderer(0, [][]byte{jpegBuf.Bytes()}, asciiOptions{
brightnessThreshold: 40,
charset: "standard",
}, newRenderCache(1<<20))
var wg sync.WaitGroup
results := make([][]byte, 32)
for i := range results {
wg.Add(1)
go func(i int) {
defer wg.Done()
ascii, err := r.render(0, 20, 10, false, colorTierTrueColor)
if err != nil {
t.Error(err)
return
}
results[i] = ascii
}(i)
}
wg.Wait()
for i, got := range results {
if !bytes.Equal(got, results[0]) {
t.Fatalf("result %d differs from result 0", i)
} }
if tr.hasReachedLimits("1.2.3.4", 2, 100) {
t.Error("should be cleared")
} }
} }
func TestClampDimension(t *testing.T) { func TestRenderCacheEvictsByBytes(t *testing.T) {
if clampDimension(0, 100) != 1 { key := func(index int) cacheKey { return cacheKey{index: index} }
t.Error("floor") budget := 3 * entryCost(key(0), bytes.Repeat([]byte("x"), 1000))
c := newRenderCache(budget)
for i := range 5 {
c.put(key(i), bytes.Repeat([]byte("x"), 1000))
} }
if clampDimension(500, 100) != 100 { if c.size.Load() > budget {
t.Error("ceil") t.Errorf("size %d exceeds budget %d", c.size.Load(), budget)
} }
if clampDimension(50, 100) != 50 { if _, ok := c.get(key(0)); ok {
t.Error("passthrough") t.Error("oldest entry should have been evicted")
}
if _, ok := c.get(key(4)); !ok {
t.Error("newest entry should be cached")
}
}
func TestRenderCacheDisabled(t *testing.T) {
c := newRenderCache(0)
if c != nil {
t.Fatal("zero budget should disable the cache")
}
c.put(cacheKey{}, []byte("v"))
if _, ok := c.get(cacheKey{}); ok {
t.Error("nil cache should never hit")
}
}
func TestRenderCacheRejectsOversizedEntry(t *testing.T) {
c := newRenderCache(256)
c.put(cacheKey{}, bytes.Repeat([]byte("x"), 10_000))
if _, ok := c.get(cacheKey{}); ok {
t.Error("entry larger than budget should not be cached")
}
if c.size.Load() != 0 {
t.Errorf("size = %d, want 0", c.size.Load())
}
}
func TestConnectionTracker(t *testing.T) {
tr := newConnectionTracker()
if _, _, ok := tr.tryAcquire("1.2.3.4", 2, 100); !ok {
t.Fatal("first acquire failed")
}
if _, _, ok := tr.tryAcquire("1.2.3.4", 2, 100); !ok {
t.Fatal("second acquire failed")
}
if _, _, ok := tr.tryAcquire("1.2.3.4", 2, 100); ok {
t.Error("expected per-ip limit rejection")
}
tr.release("1.2.3.4")
tr.release("1.2.3.4")
if tr.totalCount() != 0 {
t.Errorf("total = %d", tr.totalCount())
}
if _, _, ok := tr.tryAcquire("1.2.3.4", 2, 100); !ok {
t.Error("limit should be cleared")
}
}
func TestClampTermSize(t *testing.T) {
w, h := clampTermSize(1000, 500, 512, 65536, 4)
if w < 1 || h < 1 || w > 512 || h > 512 || w*h > 65536 {
t.Fatalf("clamped size = %dx%d", w, h)
}
if w%4 != 0 || h%4 != 0 {
t.Fatalf("size is not quantized: %dx%d", w, h)
}
w, h = clampTermSize(3, 2, 100, 100, 4)
if w != 3 || h != 2 {
t.Fatalf("small size = %dx%d", w, h)
} }
} }
+6 -2
View File
@@ -28,7 +28,9 @@ type Config struct {
MaxAuthAttempts int MaxAuthAttempts int
HandshakeTimeout time.Duration HandshakeTimeout time.Duration
MaxDimension int MaxDimension int
FrameResolution int MaxTerminalCells int
SessionTimeout time.Duration
RenderCacheMB int
BrightnessThreshold int BrightnessThreshold int
Charset string Charset string
Invert bool Invert bool
@@ -126,7 +128,9 @@ func loadConfig() Config {
MaxAuthAttempts: envInt("MAX_AUTH_ATTEMPTS", 6, 1, maxInt), MaxAuthAttempts: envInt("MAX_AUTH_ATTEMPTS", 6, 1, maxInt),
HandshakeTimeout: envDurationMs("HANDSHAKE_TIMEOUT", 10*time.Second), HandshakeTimeout: envDurationMs("HANDSHAKE_TIMEOUT", 10*time.Second),
MaxDimension: envInt("MAX_DIMENSION", 512, 1, 4096), MaxDimension: envInt("MAX_DIMENSION", 512, 1, 4096),
FrameResolution: envInt("FRAME_RESOLUTION", 360, 16, 1080), MaxTerminalCells: envInt("MAX_TERMINAL_CELLS", 500*512, 1, maxInt),
SessionTimeout: envDurationMs("SESSION_TIMEOUT", 10*time.Minute),
RenderCacheMB: envInt("RENDER_CACHE_MB", 256, 0, maxInt),
BrightnessThreshold: envInt("BRIGHTNESS_THRESHOLD", 40, 0, 100), BrightnessThreshold: envInt("BRIGHTNESS_THRESHOLD", 40, 0, 100),
Charset: envString("CHARSET", "detailed"), Charset: envString("CHARSET", "detailed"),
Invert: envBool("INVERT", false), Invert: envBool("INVERT", false),
+93 -21
View File
@@ -6,6 +6,7 @@ import (
"fmt" "fmt"
"math" "math"
"os" "os"
"sync"
) )
// .tsf container, little-endian: "TSFR" | version uint16 | fps float64 | // .tsf container, little-endian: "TSFR" | version uint16 | fps float64 |
@@ -13,9 +14,60 @@ import (
const ( const (
tsfMagic = "TSFR" tsfMagic = "TSFR"
tsfVersion = 1 tsfVersion = 1
maxTSFFPS = 240
maxTSFFrameCount = 10_000_000
) )
type frameFile struct {
data []byte
cleanup func() error
once sync.Once
err error
}
func (f *frameFile) Close() error {
if f == nil {
return nil
}
f.once.Do(func() {
if f.cleanup != nil {
f.err = f.cleanup()
}
f.data = nil
})
return f.err
}
var frameFileOwners sync.Map // map[*FramesContainer]*frameFile
func (data *FramesContainer) Close() error {
if data == nil {
return nil
}
owner, ok := frameFileOwners.LoadAndDelete(data)
if !ok {
return nil
}
data.ColorFrames = nil
return owner.(*frameFile).Close()
}
func writeTSF(output string, data *FramesContainer) error { func writeTSF(output string, data *FramesContainer) error {
if data == nil {
return fmt.Errorf("cannot write nil frames container")
}
if math.IsNaN(data.FPS) || math.IsInf(data.FPS, 0) || data.FPS <= 0 || data.FPS > maxTSFFPS {
return fmt.Errorf("cannot write .tsf: fps must be finite, positive, and at most %d", maxTSFFPS)
}
if len(data.ColorFrames) > maxTSFFrameCount || uint64(len(data.ColorFrames)) > math.MaxUint32 {
return fmt.Errorf("cannot write .tsf: frame count exceeds limit")
}
for i, frame := range data.ColorFrames {
if uint64(len(frame)) > math.MaxUint32 {
return fmt.Errorf("cannot write .tsf: frame %d length exceeds uint32", i)
}
}
f, err := os.Create(output) f, err := os.Create(output)
if err != nil { if err != nil {
return err return err
@@ -48,10 +100,17 @@ func writeTSF(output string, data *FramesContainer) error {
} }
func loadTSF(filename string) (*FramesContainer, error) { func loadTSF(filename string) (*FramesContainer, error) {
raw, err := os.ReadFile(filename) file, err := readFrameFile(filename)
if err != nil { if err != nil {
return nil, err return nil, err
} }
owned := false
defer func() {
if !owned {
_ = file.Close()
}
}()
raw := file.data
invalid := func() error { invalid := func() error {
return fmt.Errorf("invalid frames file %q: corrupt .tsf container", filename) return fmt.Errorf("invalid frames file %q: corrupt .tsf container", filename)
} }
@@ -65,27 +124,40 @@ func loadTSF(filename string) (*FramesContainer, error) {
} }
fps := math.Float64frombits(binary.LittleEndian.Uint64(raw[6:])) fps := math.Float64frombits(binary.LittleEndian.Uint64(raw[6:]))
count := binary.LittleEndian.Uint32(raw[14:]) count := binary.LittleEndian.Uint32(raw[14:])
if math.IsNaN(fps) || math.IsInf(fps, 0) || fps <= 0 || fps > maxTSFFPS {
colorFrames := make([][]byte, 0, count)
off := 18
for range count {
if off+4 > len(raw) {
return nil, invalid()
}
n := int(binary.LittleEndian.Uint32(raw[off:]))
off += 4
if off+n > len(raw) {
return nil, invalid()
}
colorFrames = append(colorFrames, raw[off:off+n])
off += n
}
if len(colorFrames) == 0 || fps <= 0 {
return nil, fmt.Errorf( return nil, fmt.Errorf(
"invalid frames file %q: expected non-empty frames and a positive fps", "invalid frames file %q: fps must be finite, greater than 0, and at most %d",
filename, filename, maxTSFFPS,
) )
} }
return &FramesContainer{ColorFrames: colorFrames, FPS: fps}, nil if count == 0 {
return nil, fmt.Errorf("invalid frames file %q: expected non-empty frames", filename)
}
if count > maxTSFFrameCount || uint64(count) > uint64((len(raw)-18)/4) {
return nil, invalid()
}
colorFrames := make([][]byte, 0, int(count))
off := 18
for range count {
if len(raw)-off < 4 {
return nil, invalid()
}
n := uint64(binary.LittleEndian.Uint32(raw[off:]))
off += 4
if n > uint64(len(raw)-off) {
return nil, invalid()
}
nativeLen := int(n)
colorFrames = append(colorFrames, raw[off:off+nativeLen])
off += nativeLen
}
if off != len(raw) {
return nil, invalid()
}
data := &FramesContainer{ColorFrames: colorFrames, FPS: fps}
frameFileOwners.Store(data, file)
owned = true
return data, nil
} }
+82
View File
@@ -0,0 +1,82 @@
package main
import (
"encoding/binary"
"math"
"os"
"path/filepath"
"strings"
"testing"
)
func tsfHeader(fps float64, count uint32) []byte {
raw := make([]byte, 18)
copy(raw, tsfMagic)
binary.LittleEndian.PutUint16(raw[4:], tsfVersion)
binary.LittleEndian.PutUint64(raw[6:], math.Float64bits(fps))
binary.LittleEndian.PutUint32(raw[14:], count)
return raw
}
func writeRawTSF(t *testing.T, raw []byte) string {
t.Helper()
path := filepath.Join(t.TempDir(), "frames.tsf")
if err := os.WriteFile(path, raw, 0o644); err != nil {
t.Fatal(err)
}
return path
}
func TestTSFRejectsInvalidFPS(t *testing.T) {
for _, fps := range []float64{math.NaN(), math.Inf(1), math.Inf(-1), -1, 0, 240.01} {
raw := append(tsfHeader(fps, 1), 0, 0, 0, 0)
if _, err := loadTSF(writeRawTSF(t, raw)); err == nil {
t.Errorf("loadTSF accepted fps %v", fps)
}
}
}
func TestTSFRejectsImpossibleCountsAndLengths(t *testing.T) {
if _, err := loadTSF(writeRawTSF(t, tsfHeader(30, math.MaxUint32))); err == nil {
t.Fatal("loadTSF accepted impossible frame count")
}
raw := append(tsfHeader(30, 1), 0xff, 0xff, 0xff, 0xff)
if _, err := loadTSF(writeRawTSF(t, raw)); err == nil {
t.Fatal("loadTSF accepted overflowing frame length")
}
}
func TestTSFCloseReleasesOwnedFrames(t *testing.T) {
path := filepath.Join(t.TempDir(), "frames.tsf")
if err := writeTSF(path, &FramesContainer{FPS: 30, ColorFrames: [][]byte{{1, 2, 3}}}); err != nil {
t.Fatal(err)
}
frames, err := loadTSF(path)
if err != nil {
t.Fatal(err)
}
if got := frames.ColorFrames[0]; len(got) != 3 || got[0] != 1 {
t.Fatalf("unexpected zero-copy frame data: %v", got)
}
if err := frames.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if frames.ColorFrames != nil {
t.Fatal("Close retained references to released frame data")
}
if err := frames.Close(); err != nil {
t.Fatalf("second Close: %v", err)
}
}
func TestTSFWriteRejectsInvalidHeaderValuesBeforeCreate(t *testing.T) {
path := filepath.Join(t.TempDir(), "frames.tsf")
err := writeTSF(path, &FramesContainer{FPS: math.NaN(), ColorFrames: [][]byte{{1}}})
if err == nil || !strings.Contains(err.Error(), "fps") {
t.Fatalf("writeTSF error = %v", err)
}
if _, err := os.Stat(path); !os.IsNotExist(err) {
t.Fatalf("invalid write created output: %v", err)
}
}
+319 -83
View File
@@ -3,12 +3,15 @@ package main
import ( import (
"bytes" "bytes"
"container/list" "container/list"
"fmt"
"image" "image"
"image/color" "image/color"
"image/jpeg" "image/jpeg"
"strconv"
"strings" "strings"
"sync" "sync"
"sync/atomic"
"time"
"unicode/utf8"
"golang.org/x/image/draw" "golang.org/x/image/draw"
) )
@@ -62,22 +65,52 @@ func resolveCharset(charset string) string {
return charset return charset
} }
func resizeFrame(frame []byte, width, height int, keepAspectRatio bool, tier colorTier) (draw.Image, error) { const maxPooledBuffer = 4 << 20
var pixPool sync.Pool
func getPixBuf(n int) []byte {
if v := pixPool.Get(); v != nil {
if b := *v.(*[]byte); cap(b) >= n {
return b[:n]
}
}
return make([]byte, n)
}
func putPixBuf(b []byte) {
if cap(b) <= maxPooledBuffer {
pixPool.Put(&b)
}
}
var outPool sync.Pool
func getOutBuf(capacity int) []byte {
if v := outPool.Get(); v != nil {
if b := *v.(*[]byte); cap(b) >= capacity {
return b[:0]
}
}
return make([]byte, 0, capacity)
}
func putOutBuf(b []byte) {
if cap(b) <= maxPooledBuffer {
outPool.Put(&b)
}
}
func resizeFrame(frame, pix []byte, width, height int, keepAspectRatio bool) (*image.RGBA, error) {
src, err := jpeg.Decode(bytes.NewReader(frame)) src, err := jpeg.Decode(bytes.NewReader(frame))
if err != nil { if err != nil {
return nil, err return nil, err
} }
var dst draw.Image rect := image.Rect(0, 0, width, height)
var bg color.Color
if tier == colorTierNone { dst := &image.RGBA{Pix: pix[:4*width*height], Stride: 4 * width, Rect: rect}
dst = image.NewGray(image.Rect(0, 0, width, height))
bg = color.Gray{0}
} else {
dst = image.NewNRGBA(image.Rect(0, 0, width, height))
bg = color.Black
}
if keepAspectRatio { if keepAspectRatio {
draw.Draw(dst, dst.Bounds(), image.NewUniform(bg), image.Point{}, draw.Src) draw.Draw(dst, dst.Bounds(), image.NewUniform(color.Black), image.Point{}, draw.Src)
sb := src.Bounds() sb := src.Bounds()
sw, sh := sb.Dx(), sb.Dy() sw, sh := sb.Dx(), sb.Dy()
scale := min(float64(width)/float64(sw), float64(height)/float64(sh)) scale := min(float64(width)/float64(sw), float64(height)/float64(sh))
@@ -98,6 +131,15 @@ type asciiOptions struct {
invert bool invert bool
} }
func buildRampLUT(ramp []rune, options asciiOptions) *[101][]byte {
var lut [101][]byte
for b := range lut {
index := rampIndex(b, options.brightnessThreshold, len(ramp), options.invert)
lut[b] = utf8.AppendRune(nil, ramp[index])
}
return &lut
}
func rampIndex(brightness, threshold, total int, invert bool) int { func rampIndex(brightness, threshold, total int, invert bool) int {
var index int var index int
if brightness < threshold { if brightness < threshold {
@@ -114,27 +156,38 @@ func rampIndex(brightness, threshold, total int, invert bool) int {
return index return index
} }
func frameToAscii(pixels []byte, options asciiOptions) string { func frameToAscii(img *image.RGBA, rampLUT *[101][]byte) []byte {
ramp := []rune(resolveCharset(options.charset)) pix := img.Pix
total := len(ramp) maxCharBytes := 1
var b strings.Builder for _, char := range rampLUT {
for _, p := range pixels { maxCharBytes = max(maxCharBytes, len(char))
brightness := int(p) * 100 / 255
index := rampIndex(brightness, options.brightnessThreshold, total, options.invert)
b.WriteRune(ramp[index])
} }
return b.String() buf := getOutBuf(len(pix) / 4 * maxCharBytes)
for o := 0; o < len(pix); o += 4 {
brightness := (int(pix[o])*299 + int(pix[o+1])*587 + int(pix[o+2])*114) / 255 / 10
buf = append(buf, rampLUT[brightness]...)
}
output := bytes.Clone(buf)
putOutBuf(buf)
return output
} }
const ansiReset = "\x1b[0m" const ansiReset = "\x1b[0m"
var ansi256Levels = [6]int{0, 95, 135, 175, 215, 255} var ansi256Levels = [6]int{0, 95, 135, 175, 215, 255}
func quantize256(r, g, b uint8) int { var decimal = func() (t [256]string) {
toLevel := func(v uint8) int { for i := range t {
t[i] = strconv.Itoa(i)
}
return
}()
var ansi256Cube = func() (t [256]uint8) {
for v := range t {
best, bestDist := 0, 1<<30 best, bestDist := 0, 1<<30
for i, l := range ansi256Levels { for i, l := range ansi256Levels {
d := int(v) - l d := v - l
if d < 0 { if d < 0 {
d = -d d = -d
} }
@@ -142,105 +195,288 @@ func quantize256(r, g, b uint8) int {
bestDist, best = d, i bestDist, best = d, i
} }
} }
return best t[v] = uint8(best)
} }
return 16 + 36*toLevel(r) + 6*toLevel(g) + toLevel(b) return
}()
func quantize256(r, g, b uint8) int {
return 16 + 36*int(ansi256Cube[r]) + 6*int(ansi256Cube[g]) + int(ansi256Cube[b])
} }
func frameToAnsi(img *image.NRGBA, options asciiOptions, tier colorTier) string { func appendColor(buf []byte, r, g, b uint8, tier colorTier) []byte {
ramp := []rune(resolveCharset(options.charset)) if tier == colorTierTrueColor {
total := len(ramp) buf = append(buf, "\x1b[38;2;"...)
buf = append(buf, decimal[r]...)
buf = append(buf, ';')
buf = append(buf, decimal[g]...)
buf = append(buf, ';')
buf = append(buf, decimal[b]...)
} else {
buf = append(buf, "\x1b[38;5;"...)
buf = append(buf, decimal[quantize256(r, g, b)]...)
}
return append(buf, 'm')
}
func frameToAnsi(img *image.RGBA, rampLUT *[101][]byte, tier colorTier) []byte {
bounds := img.Bounds() bounds := img.Bounds()
var b strings.Builder bytesPerCell := 11
if tier == colorTierTrueColor {
bytesPerCell = 16
}
buf := getOutBuf(bounds.Dx() * bounds.Dy() * bytesPerCell)
var lastR, lastG, lastB uint8 var lastR, lastG, lastB uint8
last256 := -1
first := true first := true
for y := bounds.Min.Y; y < bounds.Max.Y; y++ { for y := bounds.Min.Y; y < bounds.Max.Y; y++ {
o := img.PixOffset(bounds.Min.X, y)
for x := bounds.Min.X; x < bounds.Max.X; x++ { for x := bounds.Min.X; x < bounds.Max.X; x++ {
o := img.PixOffset(x, y)
r, g, bl := img.Pix[o], img.Pix[o+1], img.Pix[o+2] r, g, bl := img.Pix[o], img.Pix[o+1], img.Pix[o+2]
brightness := (int(r)*299 + int(g)*587 + int(bl)*114) / 255 / 10 brightness := (int(r)*299 + int(g)*587 + int(bl)*114) / 255 / 10
index := rampIndex(brightness, options.brightnessThreshold, total, options.invert) colorChanged := first || r != lastR || g != lastG || bl != lastB
if first || r != lastR || g != lastG || bl != lastB { if tier == colorTier256 {
if tier == colorTierTrueColor { index := quantize256(r, g, bl)
fmt.Fprintf(&b, "\x1b[38;2;%d;%d;%dm", r, g, bl) colorChanged = first || index != last256
} else { last256 = index
fmt.Fprintf(&b, "\x1b[38;5;%dm", quantize256(r, g, bl))
} }
if colorChanged {
buf = appendColor(buf, r, g, bl, tier)
lastR, lastG, lastB = r, g, bl lastR, lastG, lastB = r, g, bl
first = false first = false
} }
b.WriteRune(ramp[index]) buf = append(buf, rampLUT[brightness]...)
o += 4
} }
if y < bounds.Max.Y-1 { if y < bounds.Max.Y-1 {
b.WriteString(ansiReset + "\r\n") buf = append(buf, "\r\n"...)
first = true
} }
} }
b.WriteString(ansiReset) buf = append(buf, ansiReset...)
return b.String() output := bytes.Clone(buf)
putOutBuf(buf)
return output
} }
type FrameRenderer struct { type cacheKey struct {
colorFrames [][]byte setID int
options asciiOptions index int
maxEntries int width int
height int
keepAspectRatio bool
tier colorTier
}
type renderCacheShard struct {
mu sync.Mutex mu sync.Mutex
cache map[string]*list.Element maxBytes int64
size int64
entries map[cacheKey]*list.Element
order *list.List order *list.List
} }
type cacheEntry struct { type renderCache struct {
key string shards []renderCacheShard
ascii string size atomic.Int64
hits atomic.Uint64
misses atomic.Uint64
evictions atomic.Uint64
rejections atomic.Uint64
renders atomic.Uint64
renderNs atomic.Uint64
} }
func newFrameRenderer(colorFrames [][]byte, options asciiOptions) *FrameRenderer { type cacheEntry struct {
return &FrameRenderer{ key cacheKey
colorFrames: colorFrames, ascii []byte
options: options, cost int64
maxEntries: 4096, }
cache: make(map[string]*list.Element),
func entryCost(_ cacheKey, ascii []byte) int64 {
return int64(cap(ascii)) + 160
}
func newRenderCache(maxBytes int64) *renderCache {
if maxBytes <= 0 {
return nil
}
shardCount := int(min(int64(16), max(int64(1), maxBytes/(1<<20))))
cache := &renderCache{shards: make([]renderCacheShard, shardCount)}
for i := range cache.shards {
cache.shards[i] = renderCacheShard{
maxBytes: maxBytes / int64(shardCount),
entries: make(map[cacheKey]*list.Element),
order: list.New(), order: list.New(),
} }
}
return cache
} }
func (r *FrameRenderer) render(index, width, height int, keepAspectRatio bool, tier colorTier) (string, error) { func (c *renderCache) shard(key cacheKey) *renderCacheShard {
key := fmt.Sprintf("%d:%dx%d:%t:%d", index, width, height, keepAspectRatio, tier) hash := uint64(key.setID)*0x9e3779b185ebca87 ^ uint64(key.index)*0xc2b2ae3d27d4eb4f
hash ^= uint64(key.width)<<32 | uint64(uint32(key.height))
hash ^= uint64(key.tier)<<1 | uint64(boolToInt(key.keepAspectRatio))
return &c.shards[hash%uint64(len(c.shards))]
}
r.mu.Lock() func boolToInt(value bool) int {
if el, ok := r.cache[key]; ok { if value {
r.order.MoveToBack(el) return 1
ascii := el.Value.(*cacheEntry).ascii }
r.mu.Unlock() return 0
}
func (c *renderCache) get(key cacheKey) ([]byte, bool) {
if c == nil {
return nil, false
}
shard := c.shard(key)
shard.mu.Lock()
defer shard.mu.Unlock()
el, ok := shard.entries[key]
if !ok {
c.misses.Add(1)
return nil, false
}
c.hits.Add(1)
shard.order.MoveToBack(el)
return el.Value.(*cacheEntry).ascii, true
}
func (c *renderCache) put(key cacheKey, ascii []byte) {
if c == nil {
return
}
shard := c.shard(key)
cost := entryCost(key, ascii)
if cost > shard.maxBytes {
c.rejections.Add(1)
return
}
shard.mu.Lock()
defer shard.mu.Unlock()
if _, ok := shard.entries[key]; ok {
return
}
shard.entries[key] = shard.order.PushBack(&cacheEntry{key: key, ascii: ascii, cost: cost})
shard.size += cost
c.size.Add(cost)
for shard.size > shard.maxBytes {
oldest := shard.order.Front()
shard.order.Remove(oldest)
evicted := oldest.Value.(*cacheEntry)
delete(shard.entries, evicted.key)
shard.size -= evicted.cost
c.size.Add(-evicted.cost)
c.evictions.Add(1)
}
}
type renderCacheStats struct {
SizeBytes int64
Hits uint64
Misses uint64
Evictions uint64
Rejections uint64
Renders uint64
RenderTime time.Duration
}
func (c *renderCache) stats() renderCacheStats {
if c == nil {
return renderCacheStats{}
}
return renderCacheStats{
SizeBytes: c.size.Load(),
Hits: c.hits.Load(),
Misses: c.misses.Load(),
Evictions: c.evictions.Load(),
Rejections: c.rejections.Load(),
Renders: c.renders.Load(),
RenderTime: time.Duration(c.renderNs.Load()),
}
}
type FrameRenderer struct {
setID int
colorFrames [][]byte
options asciiOptions
rampLUT *[101][]byte
cache *renderCache
inflightMu sync.Mutex
inflight map[cacheKey]*renderCall
}
type renderCall struct {
done chan struct{}
value []byte
err error
}
func newFrameRenderer(setID int, colorFrames [][]byte, options asciiOptions, cache *renderCache) *FrameRenderer {
ramp := []rune(resolveCharset(options.charset))
return &FrameRenderer{
setID: setID,
colorFrames: colorFrames,
options: options,
rampLUT: buildRampLUT(ramp, options),
cache: cache,
inflight: make(map[cacheKey]*renderCall),
}
}
func (r *FrameRenderer) render(index, width, height int, keepAspectRatio bool, tier colorTier) ([]byte, error) {
key := cacheKey{r.setID, index, width, height, keepAspectRatio, tier}
if ascii, ok := r.cache.get(key); ok {
return ascii, nil return ascii, nil
} }
r.mu.Unlock()
var ascii string r.inflightMu.Lock()
if call, ok := r.inflight[key]; ok {
r.inflightMu.Unlock()
<-call.done
return call.value, call.err
}
call := &renderCall{done: make(chan struct{})}
r.inflight[key] = call
r.inflightMu.Unlock()
defer func() {
r.inflightMu.Lock()
delete(r.inflight, key)
r.inflightMu.Unlock()
close(call.done)
}()
if ascii, ok := r.cache.get(key); ok {
call.value = ascii
return ascii, nil
}
started := time.Now()
pix := getPixBuf(4 * width * height)
img, err := resizeFrame(r.colorFrames[index], pix, width, height, keepAspectRatio)
if err != nil {
putPixBuf(pix)
call.err = err
return nil, err
}
var ascii []byte
if tier == colorTierNone { if tier == colorTierNone {
img, err := resizeFrame(r.colorFrames[index], width, height, keepAspectRatio, tier) ascii = frameToAscii(img, r.rampLUT)
if err != nil {
return "", err
}
ascii = frameToAscii(img.(*image.Gray).Pix, r.options)
} else { } else {
img, err := resizeFrame(r.colorFrames[index], width, height, keepAspectRatio, tier) ascii = frameToAnsi(img, r.rampLUT, tier)
if err != nil {
return "", err
}
ascii = frameToAnsi(img.(*image.NRGBA), r.options, tier)
} }
putPixBuf(pix)
r.mu.Lock() if r.cache != nil && cap(ascii) > len(ascii)+len(ascii)/4 {
if _, ok := r.cache[key]; !ok { ascii = bytes.Clone(ascii)
r.cache[key] = r.order.PushBack(&cacheEntry{key, ascii})
if r.order.Len() > r.maxEntries {
oldest := r.order.Front()
r.order.Remove(oldest)
delete(r.cache, oldest.Value.(*cacheEntry).key)
} }
r.cache.put(key, ascii)
if r.cache != nil {
r.cache.renders.Add(1)
r.cache.renderNs.Add(uint64(time.Since(started)))
} }
r.mu.Unlock() call.value = ascii
return ascii, nil return ascii, nil
} }
+7 -5
View File
@@ -47,18 +47,20 @@ func sanitizeN(value any, maxLength int) string {
str = fmt.Sprint(v) str = fmt.Sprint(v)
} }
var b strings.Builder var b strings.Builder
count := 0
for _, r := range str { for _, r := range str {
if count >= maxLength {
b.WriteRune('…')
return b.String()
}
if r < 0x20 || (r >= 0x7f && r <= 0x9f) { if r < 0x20 || (r >= 0x7f && r <= 0x9f) {
b.WriteRune('') b.WriteRune('')
} else { } else {
b.WriteRune(r) b.WriteRune(r)
} }
count++
} }
out := []rune(b.String()) return b.String()
if len(out) > maxLength {
return string(out[:maxLength]) + "…"
}
return string(out)
} }
func emit(level logLevel, name string, stream *os.File, args []any) { func emit(level logLevel, name string, stream *os.File, args []any) {
+59 -5
View File
@@ -6,7 +6,9 @@ import (
"os/signal" "os/signal"
"path/filepath" "path/filepath"
"runtime" "runtime"
"runtime/debug"
"sort" "sort"
"strconv"
"strings" "strings"
"sync" "sync"
"syscall" "syscall"
@@ -18,10 +20,11 @@ import (
type cliArgs struct { type cliArgs struct {
generate bool generate bool
video string video string
resolution int
} }
func parseArgs(argv []string) cliArgs { func parseArgs(argv []string) cliArgs {
var args cliArgs args := cliArgs{resolution: 512}
for i := 0; i < len(argv); i++ { for i := 0; i < len(argv); i++ {
switch argv[i] { switch argv[i] {
case "--generate", "-g": case "--generate", "-g":
@@ -31,6 +34,13 @@ func parseArgs(argv []string) cliArgs {
i++ i++
args.video = argv[i] args.video = argv[i]
} }
case "--resolution", "-r":
if i+1 < len(argv) {
i++
if n, err := strconv.Atoi(argv[i]); err == nil {
args.resolution = max(n, 16)
}
}
} }
} }
return args return args
@@ -52,7 +62,7 @@ func resolveVideoPath(explicitPath string) string {
return "" return ""
} }
func generateFrames(config Config, framesDir, videoArg string) { func generateFrames(framesDir, videoArg string, resolution int) {
if videoArg == "" { if videoArg == "" {
fail("No source video given. Pass --video <path>.") fail("No source video given. Pass --video <path>.")
} }
@@ -68,11 +78,13 @@ func generateFrames(config Config, framesDir, videoArg string) {
output := filepath.Join(framesDir, base+".tsf") output := filepath.Join(framesDir, base+".tsf")
logInfo(fmt.Sprintf("Generating frames from %q -> %s", videoPath, output)) logInfo(fmt.Sprintf("Generating frames from %q -> %s", videoPath, output))
if err := processVideo(videoPath, output, config.FrameResolution); err != nil { if err := processVideo(videoPath, output, resolution); err != nil {
fail(fmt.Sprintf("Failed to generate frames from %q: %s", videoPath, err.Error())) fail(fmt.Sprintf("Failed to generate frames from %q: %s", videoPath, err.Error()))
} }
} }
const frameDataWarnBytes = 2 << 30
func loadAllFrames(framesDir string) []*FramesContainer { func loadAllFrames(framesDir string) []*FramesContainer {
entries, err := os.ReadDir(framesDir) entries, err := os.ReadDir(framesDir)
var files []string var files []string
@@ -91,6 +103,22 @@ func loadAllFrames(framesDir string) []*FramesContainer {
framesDir, framesDir,
)) ))
} }
var totalBytes int64
for _, file := range files {
info, err := os.Stat(filepath.Join(framesDir, file))
if err != nil {
fail(err.Error())
}
totalBytes += info.Size()
}
if totalBytes > frameDataWarnBytes {
logWarn(fmt.Sprintf(
"Frame data is %.1f MB of mapped memory; make sure the container memory limit leaves headroom",
float64(totalBytes)/(1<<20),
))
} else {
logInfo(fmt.Sprintf("Frame data: %.1f MB", float64(totalBytes)/(1<<20)))
}
concurrency := min(len(files), max(1, min(runtime.NumCPU(), 4))) concurrency := min(len(files), max(1, min(runtime.NumCPU(), 4)))
@@ -143,9 +171,33 @@ func loadAllFrames(framesDir string) []*FramesContainer {
return results return results
} }
func applyMemoryLimit() {
if os.Getenv("GOMEMLIMIT") != "" {
return
}
for _, path := range []string{
"/sys/fs/cgroup/memory.max",
"/sys/fs/cgroup/memory/memory.limit_in_bytes",
} {
raw, err := os.ReadFile(path)
if err != nil {
continue
}
n, err := strconv.ParseInt(strings.TrimSpace(string(raw)), 10, 64)
if err != nil || n <= 0 || n > 1<<48 {
return
}
limit := n * 9 / 10
debug.SetMemoryLimit(limit)
logInfo(fmt.Sprintf("Memory limit set to %d MB (90%% of cgroup limit)", limit>>20))
return
}
}
func main() { func main() {
_ = godotenv.Load() _ = godotenv.Load()
logThreshold = resolveThreshold() logThreshold = resolveThreshold()
applyMemoryLimit()
config := loadConfig() config := loadConfig()
args := parseArgs(os.Args[1:]) args := parseArgs(os.Args[1:])
@@ -158,7 +210,7 @@ func main() {
framesDir := filepath.Join(cwd, "frames") framesDir := filepath.Join(cwd, "frames")
if args.generate { if args.generate {
generateFrames(config, framesDir, args.video) generateFrames(framesDir, args.video, args.resolution)
return return
} }
@@ -192,14 +244,16 @@ func main() {
GoodbyeText: goodbyeText, GoodbyeText: goodbyeText,
VideoSets: videoSets, VideoSets: videoSets,
}) })
defer server.Close()
sigCh := make(chan os.Signal, 1) sigCh := make(chan os.Signal, 1)
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
go func() { go func() {
sig := <-sigCh sig := <-sigCh
logInfo(fmt.Sprintf("Received %s, shutting down...", sig)) logInfo(fmt.Sprintf("Received %s, shutting down...", sig))
forceExit := time.AfterFunc(5*time.Second, func() { os.Exit(0) })
server.Close() server.Close()
time.AfterFunc(5*time.Second, func() { os.Exit(0) }) forceExit.Stop()
}() }()
if err := server.Listen(config.Host, config.Port); err != nil { if err := server.Listen(config.Host, config.Port); err != nil {
+10
View File
@@ -0,0 +1,10 @@
//go:build !unix
package main
import "os"
func readFrameFile(filename string) (*frameFile, error) {
data, err := os.ReadFile(filename)
return &frameFile{data: data}, err
}
+37
View File
@@ -0,0 +1,37 @@
//go:build unix
package main
import (
"os"
"syscall"
)
func readFrameFile(filename string) (*frameFile, error) {
f, err := os.Open(filename)
if err != nil {
return nil, err
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil, err
}
size := info.Size()
if size <= 0 || size != int64(int(size)) {
data, err := os.ReadFile(filename)
return &frameFile{data: data}, err
}
data, err := syscall.Mmap(int(f.Fd()), 0, int(size), syscall.PROT_READ, syscall.MAP_SHARED)
if err != nil {
data, err := os.ReadFile(filename)
return &frameFile{data: data}, err
}
return &frameFile{
data: data,
cleanup: func() error {
return syscall.Munmap(data)
},
}, nil
}
+127
View File
@@ -0,0 +1,127 @@
package main
import (
"bytes"
"image"
"sync"
"testing"
"golang.org/x/crypto/ssh"
)
func TestConnectionTrackerConcurrentLimit(t *testing.T) {
tracker := newConnectionTracker()
start := make(chan struct{})
var wg sync.WaitGroup
var mu sync.Mutex
accepted := make(map[string]int)
for i := range 100 {
wg.Add(1)
go func(i int) {
defer wg.Done()
<-start
ip := string(rune('a' + i%10))
if _, _, ok := tracker.tryAcquire(ip, 3, 7); ok {
mu.Lock()
accepted[ip]++
mu.Unlock()
}
}(i)
}
close(start)
wg.Wait()
total := 0
for ip, count := range accepted {
total += count
if count > 3 {
t.Fatalf("IP %q acquired %d slots", ip, count)
}
}
if total != 7 || tracker.totalCount() != 7 {
t.Fatalf("accepted=%d tracked=%d, want 7", total, tracker.totalCount())
}
for ip, count := range accepted {
for range count {
tracker.release(ip)
}
}
}
func TestSessionTrackerLimits(t *testing.T) {
tracker := newSessionTracker()
first := &ssh.ServerConn{}
second := &ssh.ServerConn{}
if !tracker.tryAcquire(first, 1, 2) {
t.Fatal("first session rejected")
}
if tracker.tryAcquire(first, 1, 2) {
t.Fatal("per-connection limit was not enforced")
}
if !tracker.tryAcquire(second, 1, 2) {
t.Fatal("second connection session rejected")
}
if tracker.tryAcquire(&ssh.ServerConn{}, 1, 2) {
t.Fatal("global session limit was not enforced")
}
tracker.release(first)
if !tracker.tryAcquire(&ssh.ServerConn{}, 1, 2) {
t.Fatal("released slot was not reusable")
}
}
func TestTermSizeDebouncesResize(t *testing.T) {
size := &termSize{}
size.set(80, 24, 512, 500*512, true)
size.set(200, 100, 512, 500*512, false)
if w, h := size.get(); w != 80 || h != 24 {
t.Fatalf("debounced size = %dx%d", w, h)
}
size.set(200, 100, 512, 500*512, true)
if w, h := size.get(); w != 200 || h != 100 {
t.Fatalf("forced size = %dx%d", w, h)
}
}
func TestAnsi256CoalescesQuantizedColors(t *testing.T) {
img := &image.RGBA{
Pix: []byte{96, 96, 96, 255, 100, 100, 100, 255},
Stride: 8,
Rect: image.Rect(0, 0, 2, 1),
}
output := frameToAnsi(img, buildRampLUT([]rune(" .#"), asciiOptions{}), colorTier256)
if count := bytes.Count(output, []byte("\x1b[38;5;")); count != 1 {
t.Fatalf("color escape count = %d, want 1: %q", count, output)
}
}
func TestAnsiDoesNotResetEachRow(t *testing.T) {
img := &image.RGBA{
Pix: []byte{100, 100, 100, 255, 100, 100, 100, 255},
Stride: 4,
Rect: image.Rect(0, 0, 1, 2),
}
output := frameToAnsi(img, buildRampLUT([]rune(" .#"), asciiOptions{}), colorTierTrueColor)
if count := bytes.Count(output, []byte(ansiReset)); count != 1 {
t.Fatalf("reset count = %d, want 1: %q", count, output)
}
}
func TestRenderCacheAccountsRetainedCapacity(t *testing.T) {
cache := newRenderCache(512)
value := make([]byte, 1, 4096)
cache.put(cacheKey{}, value)
if _, ok := cache.get(cacheKey{}); ok {
t.Fatal("cache accepted an entry whose backing allocation exceeds its budget")
}
if cache.stats().Rejections != 1 {
t.Fatalf("rejections = %d, want 1", cache.stats().Rejections)
}
}
func TestSanitizeNStopsAtLimit(t *testing.T) {
input := "ab\x00cdefghijklmnopqrstuvwxyz"
if got := sanitizeN(input, 4); got != "abc…" {
t.Fatalf("sanitizeN = %q", got)
}
}
+67
View File
@@ -0,0 +1,67 @@
package main
import (
"path/filepath"
"sync/atomic"
"testing"
)
func loadBenchSet(b *testing.B) *FramesContainer {
b.Helper()
matches, _ := filepath.Glob("../frames/*.tsf")
if len(matches) == 0 {
b.Skip("no .tsf frame set in ../frames")
}
fc, err := loadTSF(matches[0])
if err != nil {
b.Skip("failed to load frame set:", err)
}
b.Cleanup(func() { _ = fc.Close() })
return fc
}
func benchRender(b *testing.B, tier colorTier, w, h int) {
fc := loadBenchSet(b)
r := newFrameRenderer(0, fc.ColorFrames, asciiOptions{
brightnessThreshold: 40,
charset: "detailed",
}, nil)
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
if _, err := r.render(i%len(fc.ColorFrames), w, h, false, tier); err != nil {
b.Fatal(err)
}
}
}
func BenchmarkRenderTrueColor(b *testing.B) { benchRender(b, colorTierTrueColor, 120, 40) }
func BenchmarkRender256(b *testing.B) { benchRender(b, colorTier256, 120, 40) }
func BenchmarkRenderGray(b *testing.B) { benchRender(b, colorTierNone, 120, 40) }
func BenchmarkRenderTrueBig(b *testing.B) { benchRender(b, colorTierTrueColor, 240, 70) }
func BenchmarkRenderCachedParallel(b *testing.B) {
fc := loadBenchSet(b)
r := newFrameRenderer(0, fc.ColorFrames, asciiOptions{
brightnessThreshold: 40,
charset: "detailed",
}, newRenderCache(8<<20))
frame, err := r.render(0, 120, 40, false, colorTierTrueColor)
if err != nil {
b.Fatal(err)
}
b.SetBytes(int64(len(frame)))
b.ReportAllocs()
var failures atomic.Uint64
b.ResetTimer()
b.RunParallel(func(pb *testing.PB) {
for pb.Next() {
if _, err := r.render(0, 120, 40, false, colorTierTrueColor); err != nil {
failures.Add(1)
}
}
})
if failures.Load() != 0 {
b.Fatalf("render failures: %d", failures.Load())
}
}
+305 -48
View File
@@ -4,6 +4,8 @@ import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"fmt" "fmt"
"io"
"math"
"math/rand" "math/rand"
"net" "net"
"strings" "strings"
@@ -13,7 +15,84 @@ import (
"golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh"
) )
const clearScreen = "\x1b[2J\x1b[0f" const (
clearScreen = "\x1b[2J\x1b[0f"
hideCursor = "\x1b[?25l"
showCursor = "\x1b[?25h"
syncStart = "\x1b[?2026h"
syncEnd = "\x1b[?2026l"
homeCursor = "\x1b[H"
maxSessionsPerConn = 1
terminalSizeQuantum = 4
resizeDebounce = 200 * time.Millisecond
outputStallTimeout = 15 * time.Second
)
var errOutputStalled = errors.New("SSH output stalled")
func writePartsWithTimeout(
conn *ssh.ServerConn,
channel ssh.Channel,
timeout time.Duration,
parts ...string,
) error {
write := func() error {
for _, part := range parts {
if _, err := io.WriteString(channel, part); err != nil {
return err
}
}
return nil
}
if timeout <= 0 {
return write()
}
fired := make(chan struct{})
timer := time.AfterFunc(timeout, func() {
_ = conn.Close()
close(fired)
})
err := write()
if timer.Stop() {
return err
}
<-fired
return errOutputStalled
}
func writeFrameWithTimeout(
conn *ssh.ServerConn,
channel ssh.Channel,
timeout time.Duration,
prefix string,
frame []byte,
) error {
write := func() error {
if _, err := io.WriteString(channel, syncStart+prefix); err != nil {
return err
}
if _, err := channel.Write(frame); err != nil {
return err
}
_, err := io.WriteString(channel, syncEnd)
return err
}
if timeout <= 0 {
return write()
}
fired := make(chan struct{})
timer := time.AfterFunc(timeout, func() {
_ = conn.Close()
close(fired)
})
err := write()
if timer.Stop() {
return err
}
<-fired
return errOutputStalled
}
type ConnectionTracker struct { type ConnectionTracker struct {
mu sync.Mutex mu sync.Mutex
@@ -25,15 +104,18 @@ func newConnectionTracker() *ConnectionTracker {
return &ConnectionTracker{counts: make(map[string]int)} return &ConnectionTracker{counts: make(map[string]int)}
} }
func (t *ConnectionTracker) increment(ip string) int { func (t *ConnectionTracker) tryAcquire(ip string, maxPerIP, maxTotal int) (int, int, bool) {
t.mu.Lock() t.mu.Lock()
defer t.mu.Unlock() defer t.mu.Unlock()
if t.total >= maxTotal || t.counts[ip] >= maxPerIP {
return t.counts[ip], t.total, false
}
t.counts[ip]++ t.counts[ip]++
t.total++ t.total++
return t.counts[ip] return t.counts[ip], t.total, true
} }
func (t *ConnectionTracker) decrement(ip string) { func (t *ConnectionTracker) release(ip string) {
t.mu.Lock() t.mu.Lock()
defer t.mu.Unlock() defer t.mu.Unlock()
if _, ok := t.counts[ip]; !ok { if _, ok := t.counts[ip]; !ok {
@@ -54,10 +136,40 @@ func (t *ConnectionTracker) totalCount() int {
return t.total return t.total
} }
func (t *ConnectionTracker) hasReachedLimits(ip string, maxPerIP, maxTotal int) bool { type SessionTracker struct {
mu sync.Mutex
perConn map[*ssh.ServerConn]int
total int
}
func newSessionTracker() *SessionTracker {
return &SessionTracker{perConn: make(map[*ssh.ServerConn]int)}
}
func (t *SessionTracker) tryAcquire(conn *ssh.ServerConn, maxPerConn, maxTotal int) bool {
t.mu.Lock() t.mu.Lock()
defer t.mu.Unlock() defer t.mu.Unlock()
return t.total >= maxTotal || t.counts[ip] >= maxPerIP if t.total >= maxTotal || t.perConn[conn] >= maxPerConn {
return false
}
t.perConn[conn]++
t.total++
return true
}
func (t *SessionTracker) release(conn *ssh.ServerConn) {
t.mu.Lock()
defer t.mu.Unlock()
count := t.perConn[conn]
if count <= 0 {
return
}
if count == 1 {
delete(t.perConn, conn)
} else {
t.perConn[conn] = count - 1
}
t.total--
} }
type frameSet struct { type frameSet struct {
@@ -69,10 +181,16 @@ type Server struct {
config Config config Config
sshConfig *ssh.ServerConfig sshConfig *ssh.ServerConfig
sets []frameSet sets []frameSet
cache *renderCache
tracker *ConnectionTracker tracker *ConnectionTracker
sessions *SessionTracker
fakeLogin *string fakeLogin *string
goodbye *string goodbye *string
mu sync.Mutex
listener net.Listener listener net.Listener
conns map[net.Conn]struct{}
connWG sync.WaitGroup
closing bool
closeOnce sync.Once closeOnce sync.Once
} }
@@ -85,28 +203,40 @@ type ServerDeps struct {
VideoSets []*FramesContainer VideoSets []*FramesContainer
} }
func clampDimension(value, max int) int { func clampTermSize(cols, rows, maxDimension, maxCells, quantum int) (int, int) {
if value < 1 { cols = max(cols, 1)
return 1 rows = max(rows, 1)
scale := min(1.0, float64(maxDimension)/float64(cols), float64(maxDimension)/float64(rows))
area := float64(cols) * float64(rows)
if area*scale*scale > float64(maxCells) {
scale = min(scale, math.Sqrt(float64(maxCells)/area))
} }
if value > max { cols = max(1, int(math.Floor(float64(cols)*scale)))
return max rows = max(1, int(math.Floor(float64(rows)*scale)))
if quantum > 1 {
if cols >= quantum {
cols -= cols % quantum
} }
return value if rows >= quantum {
rows -= rows % quantum
}
}
return cols, rows
} }
func createServer(deps ServerDeps) *Server { func createServer(deps ServerDeps) *Server {
config := deps.Config config := deps.Config
cache := newRenderCache(int64(config.RenderCacheMB) << 20)
sets := make([]frameSet, len(deps.VideoSets)) sets := make([]frameSet, len(deps.VideoSets))
for i, data := range deps.VideoSets { for i, data := range deps.VideoSets {
sets[i] = frameSet{ sets[i] = frameSet{
data: data, data: data,
renderer: newFrameRenderer(data.ColorFrames, asciiOptions{ renderer: newFrameRenderer(i, data.ColorFrames, asciiOptions{
brightnessThreshold: config.BrightnessThreshold, brightnessThreshold: config.BrightnessThreshold,
charset: config.Charset, charset: config.Charset,
invert: config.Invert, invert: config.Invert,
}), }, cache),
} }
} }
@@ -147,9 +277,12 @@ func createServer(deps ServerDeps) *Server {
config: config, config: config,
sshConfig: sshConfig, sshConfig: sshConfig,
sets: sets, sets: sets,
cache: cache,
tracker: newConnectionTracker(), tracker: newConnectionTracker(),
sessions: newSessionTracker(),
fakeLogin: deps.FakeLoginText, fakeLogin: deps.FakeLoginText,
goodbye: deps.GoodbyeText, goodbye: deps.GoodbyeText,
conns: make(map[net.Conn]struct{}),
} }
} }
@@ -166,7 +299,14 @@ func (s *Server) Listen(host string, port int) error {
if err != nil { if err != nil {
return err return err
} }
s.mu.Lock()
if s.closing {
s.mu.Unlock()
_ = listener.Close()
return nil
}
s.listener = listener s.listener = listener
s.mu.Unlock()
logInfo(fmt.Sprintf("TrollSSH listening on %s:%d", host, port)) logInfo(fmt.Sprintf("TrollSSH listening on %s:%d", host, port))
for { for {
conn, err := listener.Accept() conn, err := listener.Accept()
@@ -176,29 +316,68 @@ func (s *Server) Listen(host string, port int) error {
} }
return err return err
} }
go s.handleConn(conn) ip := hostOnly(conn.RemoteAddr().String())
activeForIP, total, ok := s.tracker.tryAcquire(ip, s.config.MaxConnections, s.config.MaxTotalConnections)
if !ok {
_ = conn.Close()
logWarn("Connection rejected (limit reached) from", ip)
continue
}
s.mu.Lock()
if s.closing {
s.mu.Unlock()
s.tracker.release(ip)
_ = conn.Close()
continue
}
s.conns[conn] = struct{}{}
s.connWG.Add(1)
s.mu.Unlock()
go s.handleConn(conn, ip, activeForIP, total)
} }
} }
func (s *Server) Close() { func (s *Server) Close() {
s.closeOnce.Do(func() { s.closeOnce.Do(func() {
if s.listener != nil { s.mu.Lock()
_ = s.listener.Close() s.closing = true
listener := s.listener
conns := make([]net.Conn, 0, len(s.conns))
for conn := range s.conns {
conns = append(conns, conn)
}
s.mu.Unlock()
if listener != nil {
_ = listener.Close()
}
for _, conn := range conns {
_ = conn.Close()
}
s.connWG.Wait()
stats := s.cache.stats()
if stats.Hits+stats.Misses > 0 {
logInfo(fmt.Sprintf(
"Render cache: size=%.1fMB hits=%d misses=%d evictions=%d rejected=%d renders=%d render_time=%s",
float64(stats.SizeBytes)/(1<<20), stats.Hits, stats.Misses, stats.Evictions,
stats.Rejections, stats.Renders, stats.RenderTime,
))
}
for _, set := range s.sets {
if err := set.data.Close(); err != nil {
logWarn("Failed to release frame set", set.data.Name, sanitize(err.Error()))
}
} }
}) })
} }
func (s *Server) handleConn(conn net.Conn) { func (s *Server) handleConn(conn net.Conn, ip string, activeForIP, total int) {
ip := hostOnly(conn.RemoteAddr().String()) defer func() {
s.tracker.release(ip)
if s.tracker.hasReachedLimits(ip, s.config.MaxConnections, s.config.MaxTotalConnections) { s.mu.Lock()
_ = conn.Close() delete(s.conns, conn)
logWarn("Connection rejected (limit reached) from", ip) s.mu.Unlock()
return s.connWG.Done()
} }()
activeForIP := s.tracker.increment(ip)
defer s.tracker.decrement(ip)
if s.config.HandshakeTimeout > 0 { if s.config.HandshakeTimeout > 0 {
_ = conn.SetDeadline(time.Now().Add(s.config.HandshakeTimeout)) _ = conn.SetDeadline(time.Now().Add(s.config.HandshakeTimeout))
@@ -221,22 +400,40 @@ func (s *Server) handleConn(conn net.Conn) {
setIndex := rand.Intn(len(s.sets)) setIndex := rand.Intn(len(s.sets))
logInfo(fmt.Sprintf( logInfo(fmt.Sprintf(
"New connection from %s (ip=%d, total=%d) -> playing %q", "New connection from %s (ip=%d, total=%d) -> playing %q",
ip, activeForIP, s.tracker.totalCount(), s.sets[setIndex].data.Name, ip, activeForIP, total, s.sets[setIndex].data.Name,
)) ))
go ssh.DiscardRequests(reqs) go ssh.DiscardRequests(reqs)
var sessionWG sync.WaitGroup
for newChannel := range chans { for newChannel := range chans {
if newChannel.ChannelType() != "session" { if newChannel.ChannelType() != "session" {
_ = newChannel.Reject(ssh.UnknownChannelType, "unknown channel type") _ = newChannel.Reject(ssh.UnknownChannelType, "unknown channel type")
continue continue
} }
channel, requests, err := newChannel.Accept() if !s.sessions.tryAcquire(sshConn, maxSessionsPerConn, s.config.MaxTotalConnections) {
if err != nil { _ = newChannel.Reject(ssh.ResourceShortage, "session limit reached")
continue continue
} }
go s.handleSession(sshConn, channel, requests, ip, setIndex) channel, requests, err := newChannel.Accept()
if err != nil {
s.sessions.release(sshConn)
continue
} }
sessionWG.Add(1)
go func() {
defer sessionWG.Done()
defer s.sessions.release(sshConn)
var timer *time.Timer
if s.config.SessionTimeout > 0 {
timer = time.AfterFunc(s.config.SessionTimeout, func() { _ = sshConn.Close() })
defer timer.Stop()
}
s.handleSession(sshConn, channel, requests, ip, setIndex)
}()
}
_ = sshConn.Close()
sessionWG.Wait()
logInfo("Client closed connection from", ip) logInfo("Client closed connection from", ip)
} }
@@ -244,12 +441,17 @@ type termSize struct {
mu sync.Mutex mu sync.Mutex
width int width int
height int height int
updated time.Time
} }
func (t *termSize) set(w, h, maxDim int) { func (t *termSize) set(w, h, maxDimension, maxCells int, force bool) {
t.mu.Lock() t.mu.Lock()
t.width = clampDimension(w, maxDim) if !force && time.Since(t.updated) < resizeDebounce {
t.height = clampDimension(h, maxDim) t.mu.Unlock()
return
}
t.width, t.height = clampTermSize(w, h, maxDimension, maxCells, terminalSizeQuantum)
t.updated = time.Now()
t.mu.Unlock() t.mu.Unlock()
} }
@@ -296,20 +498,22 @@ func (s *Server) handleSession(
ip string, ip string,
initialSetIndex int, initialSetIndex int,
) { ) {
defer func() { _ = channel.Close() }()
size := &termSize{} size := &termSize{}
size.set(80, 24, s.config.MaxDimension) size.set(80, 24, s.config.MaxDimension, s.config.MaxTerminalCells, true)
tier := colorTierTrueColor tier := colorTierTrueColor
if s.config.ForceGrayscale { if s.config.ForceGrayscale {
tier = colorTierNone tier = colorTierNone
} }
started := false started := false
var playDone chan struct{}
for req := range requests { for req := range requests {
switch req.Type { switch req.Type {
case "pty-req": case "pty-req":
logDebug("Opening pty for session", ip) logDebug("Opening pty for session", ip)
if cols, rows, ok := parseDims(req.Payload); ok { if cols, rows, ok := parseDims(req.Payload); ok {
size.set(cols, rows, s.config.MaxDimension) size.set(cols, rows, s.config.MaxDimension, s.config.MaxTerminalCells, true)
} }
if term, ok := parsePtyTerm(req.Payload); ok { if term, ok := parsePtyTerm(req.Payload); ok {
tier = detectColorTier(term) tier = detectColorTier(term)
@@ -323,7 +527,7 @@ func (s *Server) handleSession(
if len(req.Payload) >= 8 { if len(req.Payload) >= 8 {
cols := int(binary.BigEndian.Uint32(req.Payload)) cols := int(binary.BigEndian.Uint32(req.Payload))
rows := int(binary.BigEndian.Uint32(req.Payload[4:])) rows := int(binary.BigEndian.Uint32(req.Payload[4:]))
size.set(cols, rows, s.config.MaxDimension) size.set(cols, rows, s.config.MaxDimension, s.config.MaxTerminalCells, false)
} }
if req.WantReply { if req.WantReply {
_ = req.Reply(true, nil) _ = req.Reply(true, nil)
@@ -340,14 +544,24 @@ func (s *Server) handleSession(
_ = req.Reply(true, nil) _ = req.Reply(true, nil)
if !started { if !started {
started = true started = true
go s.playVideo(sshConn, channel, size, ip, initialSetIndex, false, tier) playDone = make(chan struct{})
playTier := tier
go func(tier colorTier) {
defer close(playDone)
s.playVideo(sshConn, channel, size, ip, initialSetIndex, false, tier)
}(playTier)
} }
case "shell": case "shell":
logDebug("Opening shell for session", ip) logDebug("Opening shell for session", ip)
_ = req.Reply(true, nil) _ = req.Reply(true, nil)
if !started { if !started {
started = true started = true
go s.playVideo(sshConn, channel, size, ip, initialSetIndex, false, tier) playDone = make(chan struct{})
playTier := tier
go func(tier colorTier) {
defer close(playDone)
s.playVideo(sshConn, channel, size, ip, initialSetIndex, false, tier)
}(playTier)
} }
default: default:
if req.WantReply { if req.WantReply {
@@ -355,6 +569,10 @@ func (s *Server) handleSession(
} }
} }
} }
_ = channel.Close()
if playDone != nil {
<-playDone
}
} }
func (s *Server) pickNextSetIndex(exclude int) int { func (s *Server) pickNextSetIndex(exclude int) int {
@@ -383,9 +601,16 @@ func (s *Server) playVideo(
w, h := size.get() w, h := size.get()
logDebug(fmt.Sprintf("Terminal size %dx%d for %s", w, h, ip)) logDebug(fmt.Sprintf("Terminal size %dx%d for %s", w, h, ip))
defer func() {
_ = writePartsWithTimeout(sshConn, channel, outputStallTimeout, showCursor)
}()
if s.fakeLogin != nil { if s.fakeLogin != nil {
_, _ = channel.Write([]byte(clearScreen)) if err := writePartsWithTimeout(
_, _ = channel.Write([]byte(*s.fakeLogin)) sshConn, channel, outputStallTimeout, clearScreen, *s.fakeLogin,
); err != nil {
return
}
} }
done := make(chan struct{}) done := make(chan struct{})
@@ -429,9 +654,17 @@ func (s *Server) playVideo(
} }
}() }()
loginTimer := time.NewTimer(config.LoginDelay)
select { select {
case <-time.After(config.LoginDelay): case <-loginTimer.C:
case <-done: case <-done:
if !loginTimer.Stop() {
<-loginTimer.C
}
return
}
if err := writePartsWithTimeout(sshConn, channel, outputStallTimeout, hideCursor); err != nil {
return return
} }
@@ -444,7 +677,7 @@ func (s *Server) playVideo(
currentFrame := 0 currentFrame := 0
loopCount := 0 loopCount := 0
lastW, lastH := 0, 0
for { for {
select { select {
case <-done: case <-done:
@@ -457,6 +690,7 @@ func (s *Server) playVideo(
setIndex = (setIndex + delta + len(s.sets)) % len(s.sets) setIndex = (setIndex + delta + len(s.sets)) % len(s.sets)
current = s.sets[setIndex] current = s.sets[setIndex]
currentFrame = 0 currentFrame = 0
lastW, lastH = 0, 0
logDebug(fmt.Sprintf("%s switched to %q", ip, current.data.Name)) logDebug(fmt.Sprintf("%s switched to %q", ip, current.data.Name))
ticker.Reset(frameInterval()) ticker.Reset(frameInterval())
@@ -469,7 +703,14 @@ func (s *Server) playVideo(
return return
} }
if _, err := channel.Write([]byte(clearScreen + ascii)); err != nil { prefix := homeCursor
if w != lastW || h != lastH {
prefix = clearScreen
lastW, lastH = w, h
}
if err := writeFrameWithTimeout(
sshConn, channel, outputStallTimeout, prefix, ascii,
); err != nil {
closeSession() closeSession()
return return
} }
@@ -482,11 +723,27 @@ func (s *Server) playVideo(
currentFrame = 0 currentFrame = 0
loopCount++ loopCount++
if config.MaxLoop > 0 && loopCount >= config.MaxLoop { if config.MaxLoop > 0 && loopCount >= config.MaxLoop {
_, _ = channel.Write([]byte(clearScreen)) if err := writePartsWithTimeout(
sshConn, channel, outputStallTimeout, showCursor, clearScreen,
); err != nil {
return
}
if s.goodbye != nil { if s.goodbye != nil {
_, _ = channel.Write([]byte(*s.goodbye)) if err := writePartsWithTimeout(
sshConn, channel, outputStallTimeout, *s.goodbye,
); err != nil {
return
}
}
closeTimer := time.NewTimer(time.Second)
select {
case <-closeTimer.C:
case <-done:
if !closeTimer.Stop() {
<-closeTimer.C
}
return
} }
time.Sleep(1 * time.Second)
logInfo("Playback finished, closing session", ip) logInfo("Playback finished, closing session", ip)
_ = channel.Close() _ = channel.Close()
_ = sshConn.Close() _ = sshConn.Close()
+219 -39
View File
@@ -1,15 +1,20 @@
package main package main
import ( import (
"bufio"
"bytes" "bytes"
"encoding/binary"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"io" "io"
"math" "math"
"os"
"os/exec" "os/exec"
"path/filepath"
"strconv" "strconv"
"strings" "strings"
"time"
) )
var ( var (
@@ -17,29 +22,177 @@ var (
jpegEOI = []byte{0xff, 0xd9} jpegEOI = []byte{0xff, 0xd9}
) )
const (
maxJPEGFrameBytes = 64 << 20
maxFFmpegLogBytes = 64 << 10
)
type jpegFrameSplitter struct { type jpegFrameSplitter struct {
buffer []byte buffer []byte
scan int
inJPEG bool
} }
func (s *jpegFrameSplitter) push(chunk []byte) [][]byte { func (s *jpegFrameSplitter) push(chunk []byte, emit func([]byte) error) (int, error) {
s.buffer = append(s.buffer, chunk...) s.buffer = append(s.buffer, chunk...)
var frames [][]byte emitted := 0
for { for {
start := bytes.Index(s.buffer, jpegSOI) if !s.inJPEG {
start := bytes.Index(s.buffer[s.scan:], jpegSOI)
if start == -1 { if start == -1 {
break // Retain only a possible marker prefix spanning two reads.
if len(s.buffer) > 0 && s.buffer[len(s.buffer)-1] == jpegSOI[0] {
s.buffer = s.buffer[len(s.buffer)-1:]
} else {
s.buffer = s.buffer[:0]
} }
end := bytes.Index(s.buffer[start+len(jpegSOI):], jpegEOI) s.scan = 0
return emitted, nil
}
start += s.scan
s.buffer = s.buffer[start:]
s.scan = len(jpegSOI)
s.inJPEG = true
}
end := bytes.Index(s.buffer[s.scan:], jpegEOI)
if end == -1 { if end == -1 {
break if len(s.buffer) > maxJPEGFrameBytes {
return emitted, fmt.Errorf("JPEG frame exceeds %d MiB limit", maxJPEGFrameBytes>>20)
} }
frameEnd := start + len(jpegSOI) + end + len(jpegEOI) s.scan = max(len(jpegSOI), len(s.buffer)-1)
frame := make([]byte, frameEnd-start) return emitted, nil
copy(frame, s.buffer[start:frameEnd]) }
frames = append(frames, frame) frameEnd := s.scan + end + len(jpegEOI)
if frameEnd > maxJPEGFrameBytes {
return emitted, fmt.Errorf("JPEG frame exceeds %d MiB limit", maxJPEGFrameBytes>>20)
}
if err := emit(s.buffer[:frameEnd]); err != nil {
return emitted, err
}
emitted++
s.buffer = s.buffer[frameEnd:] s.buffer = s.buffer[frameEnd:]
s.scan = 0
s.inJPEG = false
}
}
func (s *jpegFrameSplitter) finish() error {
if s.inJPEG {
return fmt.Errorf("ffmpeg produced a truncated JPEG frame")
}
return nil
}
type boundedLog struct {
buffer bytes.Buffer
limit int
}
func (w *boundedLog) Write(p []byte) (int, error) {
n := len(p)
if remaining := w.limit - w.buffer.Len(); remaining > 0 {
_, _ = w.buffer.Write(p[:min(len(p), remaining)])
}
return n, nil
}
func (w *boundedLog) String() string {
return strings.TrimSpace(w.buffer.String())
}
type streamingTSF struct {
file *os.File
writer *bufio.Writer
path string
count uint32
}
func newStreamingTSF(output string, fps float64) (*streamingTSF, error) {
if math.IsNaN(fps) || math.IsInf(fps, 0) || fps <= 0 || fps > maxTSFFPS {
return nil, fmt.Errorf("cannot write .tsf: fps must be finite and between 0 and %d", maxTSFFPS)
}
dir := filepath.Dir(output)
f, err := os.CreateTemp(dir, "."+filepath.Base(output)+"-*.tmp")
if err != nil {
return nil, err
}
s := &streamingTSF{file: f, writer: bufio.NewWriterSize(f, 1<<20), path: f.Name()}
if err := f.Chmod(0o644); err != nil {
s.abort()
return nil, err
}
if _, err := s.writer.WriteString(tsfMagic); err != nil {
s.abort()
return nil, err
}
var hdr [14]byte
binary.LittleEndian.PutUint16(hdr[0:], tsfVersion)
binary.LittleEndian.PutUint64(hdr[2:], math.Float64bits(fps))
if _, err := s.writer.Write(hdr[:]); err != nil {
s.abort()
return nil, err
}
return s, nil
}
func (s *streamingTSF) addFrame(frame []byte) error {
if s.count >= maxTSFFrameCount {
return fmt.Errorf("too many video frames")
}
if uint64(len(frame)) > math.MaxUint32 {
return fmt.Errorf("JPEG frame is too large")
}
var size [4]byte
binary.LittleEndian.PutUint32(size[:], uint32(len(frame)))
if _, err := s.writer.Write(size[:]); err != nil {
return err
}
if _, err := s.writer.Write(frame); err != nil {
return err
}
s.count++
return nil
}
func (s *streamingTSF) commit(output string) error {
if s.count == 0 {
return fmt.Errorf("no frames were decoded from the video")
}
if err := s.writer.Flush(); err != nil {
return err
}
if _, err := s.file.Seek(14, io.SeekStart); err != nil {
return err
}
var count [4]byte
binary.LittleEndian.PutUint32(count[:], s.count)
if _, err := s.file.Write(count[:]); err != nil {
return err
}
if err := s.file.Sync(); err != nil {
return err
}
if err := s.file.Close(); err != nil {
return err
}
s.file = nil
if err := os.Rename(s.path, output); err != nil {
return err
}
s.path = ""
return nil
}
func (s *streamingTSF) abort() {
if s.file != nil {
_ = s.file.Close()
s.file = nil
}
if s.path != "" {
_ = os.Remove(s.path)
s.path = ""
} }
return frames
} }
type ffprobeOutput struct { type ffprobeOutput struct {
@@ -72,59 +225,82 @@ func parseFrameRate(rate string) float64 {
return num return num
} }
func extractFrames(path, vf, label string, maxDimension, totalFrames int) ([][]byte, error) { func extractFrames(path, vf, label string, totalFrames int, emit func([]byte) error) (int, error) {
cmd := exec.Command( cmd := exec.Command(
"ffmpeg", "-i", path, "ffmpeg", "-hide_banner", "-loglevel", "error", "-nostats",
"-i", path,
"-c:v", "mjpeg", "-c:v", "mjpeg",
"-q:v", "3", "-q:v", "3",
"-vf", vf, "-vf", vf,
"-f", "image2pipe", "-f", "image2pipe",
"pipe:1", "pipe:1",
) )
var stderr bytes.Buffer stderr := &boundedLog{limit: maxFFmpegLogBytes}
cmd.Stderr = &stderr cmd.Stderr = stderr
stdout, err := cmd.StdoutPipe() stdout, err := cmd.StdoutPipe()
if err != nil { if err != nil {
return nil, fmt.Errorf("ffmpeg failed: %w", err) return 0, fmt.Errorf("ffmpeg failed: %w", err)
} }
if err := cmd.Start(); err != nil { if err := cmd.Start(); err != nil {
return nil, fmt.Errorf("ffmpeg failed: %w", err) return 0, fmt.Errorf("ffmpeg failed: %w", err)
} }
var frames [][]byte
splitter := &jpegFrameSplitter{} splitter := &jpegFrameSplitter{}
reportProgress := func(count int) { frameCount := 0
if totalFrames > 0 { lastReport := time.Now()
pct := min(100, int(math.Round(float64(count)/float64(totalFrames)*100))) reportProgress := func(force bool) {
fmt.Printf("\rGenerating %s frames: %d/%d (%d%%)", label, count, totalFrames, pct) if !force && time.Since(lastReport) < 250*time.Millisecond {
} else { return
fmt.Printf("\rGenerating %s frames: %d", label, count)
} }
if totalFrames > 0 {
pct := min(100, int(math.Round(float64(frameCount)/float64(totalFrames)*100)))
fmt.Printf("\rGenerating %s frames: %d/%d (%d%%)", label, frameCount, totalFrames, pct)
} else {
fmt.Printf("\rGenerating %s frames: %d", label, frameCount)
}
lastReport = time.Now()
}
failStream := func(err error) (int, error) {
_ = cmd.Process.Kill()
_ = cmd.Wait()
return frameCount, err
} }
buf := make([]byte, 256*1024) buf := make([]byte, 256*1024)
for { for {
n, err := stdout.Read(buf) n, readErr := stdout.Read(buf)
if n > 0 { if n > 0 {
frames = append(frames, splitter.push(buf[:n])...) emitted, splitErr := splitter.push(buf[:n], emit)
reportProgress(len(frames)) frameCount += emitted
if splitErr != nil {
return failStream(fmt.Errorf("ffmpeg stream error: %w", splitErr))
} }
if err == io.EOF { if emitted > 0 {
reportProgress(false)
}
}
if readErr == io.EOF {
break break
} }
if err != nil { if readErr != nil {
_ = cmd.Wait() return failStream(fmt.Errorf("ffmpeg stream error: %w", readErr))
return nil, fmt.Errorf("ffmpeg stream error: %s", err.Error())
} }
} }
if err := cmd.Wait(); err != nil { if err := cmd.Wait(); err != nil {
return nil, fmt.Errorf("ffmpeg failed: %s", strings.TrimSpace(stderr.String())) if msg := stderr.String(); msg != "" {
return frameCount, fmt.Errorf("ffmpeg failed: %s", msg)
} }
if len(frames) == 0 { return frameCount, fmt.Errorf("ffmpeg failed: %w", err)
return nil, fmt.Errorf("no frames were decoded from the video")
} }
if err := splitter.finish(); err != nil {
return frameCount, err
}
if frameCount == 0 {
return 0, fmt.Errorf("no frames were decoded from the video")
}
reportProgress(true)
fmt.Println() fmt.Println()
return frames, nil return frameCount, nil
} }
func processVideo(path, output string, maxDimension int) error { func processVideo(path, output string, maxDimension int) error {
@@ -175,15 +351,19 @@ func processVideo(path, output string, maxDimension int) error {
maxDimension, maxDimension, maxDimension, maxDimension,
) )
colorFrames, err := extractFrames(path, scaleFilter, "color", maxDimension, totalFrames) outputFile, err := newStreamingTSF(output, fps)
if err != nil { if err != nil {
return err return err
} }
defer outputFile.abort()
videoData := FramesContainer{FPS: fps, ColorFrames: colorFrames} frameCount, err := extractFrames(path, scaleFilter, "color", totalFrames, outputFile.addFrame)
if err := writeTSF(output, &videoData); err != nil { if err != nil {
return err return err
} }
logInfo(fmt.Sprintf("Saved %d frames to %s", len(videoData.ColorFrames), output)) if err := outputFile.commit(output); err != nil {
return err
}
logInfo(fmt.Sprintf("Saved %d frames to %s", frameCount, output))
return nil return nil
} }
+122
View File
@@ -0,0 +1,122 @@
package main
import (
"bytes"
"os"
"path/filepath"
"testing"
)
func TestJPEGFrameSplitterAcrossChunks(t *testing.T) {
splitter := &jpegFrameSplitter{}
var frames [][]byte
emit := func(frame []byte) error {
frames = append(frames, bytes.Clone(frame))
return nil
}
chunks := [][]byte{
{0x01, 0x02, 0xff},
{0xd8, 0x10, 0xff},
{0xd9, 0xff, 0xd8, 0x20},
{0x30, 0xff},
{0xd9, 0x03},
}
for _, chunk := range chunks {
if _, err := splitter.push(chunk, emit); err != nil {
t.Fatalf("push: %v", err)
}
}
if err := splitter.finish(); err != nil {
t.Fatalf("finish: %v", err)
}
want := [][]byte{
{0xff, 0xd8, 0x10, 0xff, 0xd9},
{0xff, 0xd8, 0x20, 0x30, 0xff, 0xd9},
}
if len(frames) != len(want) {
t.Fatalf("got %d frames, want %d", len(frames), len(want))
}
for i := range want {
if !bytes.Equal(frames[i], want[i]) {
t.Errorf("frame %d = %x, want %x", i, frames[i], want[i])
}
}
}
func TestJPEGFrameSplitterRejectsTruncatedFrame(t *testing.T) {
splitter := &jpegFrameSplitter{}
if _, err := splitter.push([]byte{0xff, 0xd8, 0x01}, func([]byte) error { return nil }); err != nil {
t.Fatalf("push: %v", err)
}
if err := splitter.finish(); err == nil {
t.Fatal("finish accepted a truncated JPEG")
}
}
func TestBoundedLog(t *testing.T) {
log := &boundedLog{limit: 4}
if n, err := log.Write([]byte("abcdefgh")); err != nil || n != 8 {
t.Fatalf("Write = %d, %v", n, err)
}
if got := log.String(); got != "abcd" {
t.Fatalf("String = %q, want %q", got, "abcd")
}
}
func TestStreamingTSFCommitAndAbort(t *testing.T) {
dir := t.TempDir()
output := filepath.Join(dir, "frames.tsf")
stream, err := newStreamingTSF(output, 24)
if err != nil {
t.Fatalf("newStreamingTSF: %v", err)
}
for _, frame := range [][]byte{{1, 2, 3}, {4, 5}} {
if err := stream.addFrame(frame); err != nil {
t.Fatalf("addFrame: %v", err)
}
}
if err := stream.commit(output); err != nil {
t.Fatalf("commit: %v", err)
}
stream.abort()
got, err := loadTSF(output)
if err != nil {
t.Fatalf("loadTSF: %v", err)
}
defer func() { _ = got.Close() }()
if got.FPS != 24 || len(got.ColorFrames) != 2 || !bytes.Equal(got.ColorFrames[1], []byte{4, 5}) {
t.Fatalf("unexpected streamed TSF: %+v", got)
}
original := []byte("existing destination")
if err := os.WriteFile(output, original, 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
failed, err := newStreamingTSF(output, 24)
if err != nil {
t.Fatalf("newStreamingTSF: %v", err)
}
if err := failed.addFrame([]byte{9}); err != nil {
t.Fatalf("addFrame: %v", err)
}
failed.abort()
contents, err := os.ReadFile(output)
if err != nil {
t.Fatalf("ReadFile: %v", err)
}
if !bytes.Equal(contents, original) {
t.Fatalf("destination changed after abort: %q", contents)
}
matches, err := filepath.Glob(filepath.Join(dir, ".frames.tsf-*.tmp"))
if err != nil {
t.Fatalf("Glob: %v", err)
}
if len(matches) != 0 {
t.Fatalf("temporary files remain after abort: %v", matches)
}
}