♻️ refactor: split src into internal packages and cmd

This commit is contained in:
2026-07-16 23:51:39 +07:00
parent 2e320b03f3
commit a05d6bb7eb
28 changed files with 1355 additions and 1328 deletions
+152
View File
@@ -0,0 +1,152 @@
package config
import (
"fmt"
"os"
"strconv"
"strings"
"time"
"github.com/YuzuZensai/TrollSSH/internal/logx"
)
type PlaybackMode string
const (
PlaybackLoop PlaybackMode = "loop"
PlaybackRandom PlaybackMode = "random"
)
type Config struct {
Host string
Port int
MaxLoop int
PlaybackMode PlaybackMode
AllowUserControl bool
SwitchDebounce time.Duration
LoginDelay time.Duration
MaxConnections int
MaxTotalConnections int
MaxAuthAttempts int
HandshakeTimeout time.Duration
MaxDimension int
MaxTerminalCells int
SessionTimeout time.Duration
RenderCacheMB int
BrightnessThreshold int
Charset string
Invert bool
ForceGrayscale bool
LogCredentials bool
}
func warnInvalid(name, value string, fallback any) {
logx.Warn(fmt.Sprintf("Invalid %s=%q, using default %v", name, logx.Sanitize(value), fallback))
}
func envString(name, fallback string) string {
value := strings.TrimSpace(os.Getenv(name))
if value == "" {
return fallback
}
return value
}
func envInt(name string, fallback, min, max int) int {
raw := strings.TrimSpace(os.Getenv(name))
if raw == "" {
return fallback
}
parsed, err := strconv.Atoi(raw)
if err != nil {
warnInvalid(name, raw, fallback)
return fallback
}
if parsed < min {
return min
}
if parsed > max {
return max
}
return parsed
}
func envBool(name string, fallback bool) bool {
raw := strings.TrimSpace(os.Getenv(name))
if raw == "" {
return fallback
}
parsed, err := strconv.ParseBool(strings.ToLower(raw))
if err != nil {
warnInvalid(name, raw, fallback)
return fallback
}
return parsed
}
func envDurationMs(name string, fallback time.Duration) time.Duration {
raw := strings.TrimSpace(os.Getenv(name))
if raw == "" {
return fallback
}
ms, err := strconv.Atoi(raw)
if err != nil {
warnInvalid(name, raw, fallback)
return fallback
}
if ms < 0 {
ms = 0
}
return time.Duration(ms) * time.Millisecond
}
func envPlaybackMode(name string, fallback PlaybackMode) PlaybackMode {
raw := strings.TrimSpace(os.Getenv(name))
if raw == "" {
return fallback
}
switch PlaybackMode(strings.ToLower(raw)) {
case PlaybackLoop:
return PlaybackLoop
case PlaybackRandom:
return PlaybackRandom
}
warnInvalid(name, raw, fallback)
return fallback
}
func Load() Config {
const maxInt = int(^uint(0) >> 1)
return Config{
Host: envString("HOST", "0.0.0.0"),
Port: envInt("PORT", 22, 1, 65535),
MaxLoop: envInt("MAX_LOOP", 5, 0, maxInt),
PlaybackMode: envPlaybackMode("PLAYBACK_MODE", PlaybackLoop),
AllowUserControl: envBool("ALLOW_USER_CONTROL", true),
SwitchDebounce: envDurationMs("SWITCH_DEBOUNCE_MS", 120*time.Millisecond),
LoginDelay: envDurationMs("LOGIN_DELAY", 1500*time.Millisecond),
MaxConnections: envInt("MAX_CONNECTIONS", 10, 1, maxInt),
MaxTotalConnections: envInt("MAX_TOTAL_CONNECTIONS", 1000, 1, maxInt),
MaxAuthAttempts: envInt("MAX_AUTH_ATTEMPTS", 6, 1, maxInt),
HandshakeTimeout: envDurationMs("HANDSHAKE_TIMEOUT", 10*time.Second),
MaxDimension: envInt("MAX_DIMENSION", 512, 1, 4096),
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),
Charset: envString("CHARSET", "detailed"),
Invert: envBool("INVERT", false),
ForceGrayscale: envBool("FORCE_GRAYSCALE", false),
LogCredentials: envBool("LOG_CREDENTIALS", false),
}
}
func LoadOptionalTextFile(filePath string) (string, bool) {
data, err := os.ReadFile(filePath)
if err != nil {
return "", false
}
text := strings.ReplaceAll(string(data), "\r\n", "\n")
text = strings.ReplaceAll(text, "\n", "\r\n")
return text, true
}
+79
View File
@@ -0,0 +1,79 @@
package config
import (
"testing"
"time"
)
func TestLoadConfigDefaults(t *testing.T) {
t.Setenv("HOST", "")
t.Setenv("PORT", "")
t.Setenv("PLAYBACK_MODE", "")
t.Setenv("LOGIN_DELAY", "")
cfg := Load()
if cfg.Host != "0.0.0.0" {
t.Errorf("host = %q", cfg.Host)
}
if cfg.Port != 22 {
t.Errorf("port = %d", cfg.Port)
}
if cfg.PlaybackMode != PlaybackLoop {
t.Errorf("playbackMode = %q", cfg.PlaybackMode)
}
if cfg.Charset != "detailed" {
t.Errorf("charset = %q", cfg.Charset)
}
if cfg.LoginDelay != 1500*time.Millisecond {
t.Errorf("loginDelay = %v", cfg.LoginDelay)
}
}
func TestLoadConfigClamping(t *testing.T) {
t.Setenv("PORT", "999999")
t.Setenv("BRIGHTNESS_THRESHOLD", "-5")
cfg := Load()
if cfg.Port != 65535 {
t.Errorf("port clamp = %d", cfg.Port)
}
if cfg.BrightnessThreshold != 0 {
t.Errorf("brightness clamp = %d", cfg.BrightnessThreshold)
}
}
func TestLoadConfigInvalidFallsBack(t *testing.T) {
t.Setenv("PORT", "not-a-number")
t.Setenv("INVERT", "yes-please")
t.Setenv("PLAYBACK_MODE", "shuffle")
cfg := Load()
if cfg.Port != 22 {
t.Errorf("port = %d, want default 22", cfg.Port)
}
if cfg.Invert {
t.Error("invert should fall back to false")
}
if cfg.PlaybackMode != PlaybackLoop {
t.Errorf("playbackMode = %q, want default loop", cfg.PlaybackMode)
}
}
func TestEnvDurationMs(t *testing.T) {
t.Setenv("D", "250")
if got := envDurationMs("D", time.Second); got != 250*time.Millisecond {
t.Errorf("250 = %v, want 250ms", got)
}
t.Setenv("D", "-10")
if got := envDurationMs("D", time.Second); got != 0 {
t.Errorf("negative = %v, want 0", got)
}
t.Setenv("D", "banana")
if got := envDurationMs("D", time.Second); got != time.Second {
t.Errorf("invalid = %v, want fallback 1s", got)
}
}
func TestPlaybackModeRandom(t *testing.T) {
t.Setenv("PLAYBACK_MODE", "RaNdOm")
if Load().PlaybackMode != PlaybackRandom {
t.Error("expected random")
}
}
+89
View File
@@ -0,0 +1,89 @@
package logx
import (
"encoding/json"
"fmt"
"os"
"strings"
"time"
)
type Level int
const (
LevelDebug Level = 10
LevelInfo Level = 20
LevelWarn Level = 30
LevelError Level = 40
)
var threshold = ResolveThreshold()
func ResolveThreshold() Level {
switch strings.ToLower(strings.TrimSpace(os.Getenv("LOG_LEVEL"))) {
case "debug":
return LevelDebug
case "warn":
return LevelWarn
case "error":
return LevelError
default:
return LevelInfo
}
}
func SetThreshold(level Level) { threshold = level }
func Sanitize(value any) string {
return SanitizeN(value, 200)
}
func SanitizeN(value any, maxLength int) string {
var str string
switch v := value.(type) {
case nil:
str = ""
case string:
str = v
default:
str = fmt.Sprint(v)
}
var b strings.Builder
count := 0
for _, r := range str {
if count >= maxLength {
b.WriteRune('…')
return b.String()
}
if r < 0x20 || (r >= 0x7f && r <= 0x9f) {
b.WriteRune('')
} else {
b.WriteRune(r)
}
count++
}
return b.String()
}
func emit(level Level, name string, stream *os.File, args []any) {
if level < threshold {
return
}
parts := make([]string, len(args))
for i, a := range args {
if s, ok := a.(string); ok {
parts[i] = s
} else if b, err := json.Marshal(a); err == nil {
parts[i] = string(b)
} else {
parts[i] = fmt.Sprint(a)
}
}
ts := time.Now().UTC().Format("2006-01-02T15:04:05.000Z")
_, _ = fmt.Fprintf(stream, "[%s] %-5s %s\n", ts, strings.ToUpper(name), strings.Join(parts, " "))
}
func Debug(args ...any) { emit(LevelDebug, "debug", os.Stdout, args) }
func Info(args ...any) { emit(LevelInfo, "info", os.Stdout, args) }
func Warn(args ...any) { emit(LevelWarn, "warn", os.Stderr, args) }
func Error(args ...any) { emit(LevelError, "error", os.Stderr, args) }
+10
View File
@@ -0,0 +1,10 @@
package logx
import "testing"
func TestSanitizeNStopsAtLimit(t *testing.T) {
input := "ab\x00cdefghijklmnopqrstuvwxyz"
if got := SanitizeN(input, 4); got != "abc…" {
t.Fatalf("SanitizeN = %q", got)
}
}
+232
View File
@@ -0,0 +1,232 @@
package render
import (
"bytes"
"container/list"
"sync"
"sync/atomic"
"time"
)
type cacheKey struct {
setID int
index int
width int
height int
keepAspectRatio bool
tier ColorTier
}
type cacheShard struct {
mu sync.Mutex
maxBytes int64
size int64
entries map[cacheKey]*list.Element
order *list.List
}
type Cache struct {
shards []cacheShard
size atomic.Int64
hits atomic.Uint64
misses atomic.Uint64
evictions atomic.Uint64
rejections atomic.Uint64
renders atomic.Uint64
renderNs atomic.Uint64
}
type cacheEntry struct {
key cacheKey
ascii []byte
cost int64
}
func entryCost(_ cacheKey, ascii []byte) int64 {
return int64(cap(ascii)) + 160
}
func NewCache(maxBytes int64) *Cache {
if maxBytes <= 0 {
return nil
}
shardCount := int(min(int64(16), max(int64(1), maxBytes/(1<<20))))
cache := &Cache{shards: make([]cacheShard, shardCount)}
for i := range cache.shards {
cache.shards[i] = cacheShard{
maxBytes: maxBytes / int64(shardCount),
entries: make(map[cacheKey]*list.Element),
order: list.New(),
}
}
return cache
}
func (c *Cache) shard(key cacheKey) *cacheShard {
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))]
}
func boolToInt(value bool) int {
if value {
return 1
}
return 0
}
func (c *Cache) 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 *Cache) 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 CacheStats struct {
SizeBytes int64
Hits uint64
Misses uint64
Evictions uint64
Rejections uint64
Renders uint64
RenderTime time.Duration
}
func (c *Cache) Stats() CacheStats {
if c == nil {
return CacheStats{}
}
return CacheStats{
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 Renderer struct {
setID int
colorFrames [][]byte
options Options
rampLUT *[101][]byte
cache *Cache
inflightMu sync.Mutex
inflight map[cacheKey]*renderCall
}
type renderCall struct {
done chan struct{}
value []byte
err error
}
func NewRenderer(setID int, colorFrames [][]byte, options Options, cache *Cache) *Renderer {
ramp := []rune(resolveCharset(options.Charset))
return &Renderer{
setID: setID,
colorFrames: colorFrames,
options: options,
rampLUT: buildRampLUT(ramp, options),
cache: cache,
inflight: make(map[cacheKey]*renderCall),
}
}
func (r *Renderer) 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
}
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 {
ascii = frameToAscii(img, r.rampLUT)
} else {
ascii = frameToAnsi(img, r.rampLUT, tier)
}
putPixBuf(pix)
if r.cache != nil && cap(ascii) > len(ascii)+len(ascii)/4 {
ascii = bytes.Clone(ascii)
}
r.cache.put(key, ascii)
if r.cache != nil {
r.cache.renders.Add(1)
r.cache.renderNs.Add(uint64(time.Since(started)))
}
call.value = ascii
return ascii, nil
}
+250
View File
@@ -0,0 +1,250 @@
package render
import (
"bytes"
"image"
"image/color"
"image/jpeg"
"strconv"
"strings"
"sync"
"unicode/utf8"
"golang.org/x/image/draw"
)
type ColorTier int
const (
ColorTierNone ColorTier = iota
ColorTier256
ColorTierTrueColor
)
func DetectColorTier(term string) ColorTier {
t := strings.ToLower(strings.TrimSpace(term))
switch t {
case "", "dumb", "vt52", "vt100", "vt102", "vt220", "ansi", "linux", "cons25", "cygwin":
return ColorTierNone
}
if strings.Contains(t, "direct") || strings.Contains(t, "truecolor") {
return ColorTierTrueColor
}
if strings.Contains(t, "256color") {
return ColorTier256
}
if strings.HasPrefix(t, "screen") || strings.HasPrefix(t, "tmux") {
return ColorTier256
}
return ColorTierTrueColor
}
var charsetPresets = map[string]string{
"detailed": " .'`^\",:;Il!i><~+_-?][}{1)(|/tfjrxnuvczXYUJCLQ0OZmwqpdbkhao*#MW&8%B@$",
"standard": " .:-=+*#%@",
"simple": " .:oO#@",
"blocks": " ░▒▓█",
}
func resolveCharset(charset string) string {
if charset == "" {
return charsetPresets["detailed"]
}
if preset, ok := charsetPresets[strings.ToLower(charset)]; ok {
return preset
}
return charset
}
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))
if err != nil {
return nil, err
}
rect := image.Rect(0, 0, width, height)
dst := &image.RGBA{Pix: pix[:4*width*height], Stride: 4 * width, Rect: rect}
if keepAspectRatio {
draw.Draw(dst, dst.Bounds(), image.NewUniform(color.Black), image.Point{}, draw.Src)
sb := src.Bounds()
sw, sh := sb.Dx(), sb.Dy()
scale := min(float64(width)/float64(sw), float64(height)/float64(sh))
tw := max(1, int(float64(sw)*scale))
th := max(1, int(float64(sh)*scale))
x0 := (width - tw) / 2
y0 := (height - th) / 2
draw.ApproxBiLinear.Scale(dst, image.Rect(x0, y0, x0+tw, y0+th), src, sb, draw.Src, nil)
} else {
draw.ApproxBiLinear.Scale(dst, dst.Bounds(), src, src.Bounds(), draw.Src, nil)
}
return dst, nil
}
type Options struct {
BrightnessThreshold int
Charset string
Invert bool
}
func buildRampLUT(ramp []rune, options Options) *[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 {
var index int
if brightness < threshold {
index = 0
} else {
index = brightness * total / 100
if index > total-1 {
index = total - 1
}
}
if invert {
index = total - 1 - index
}
return index
}
func frameToAscii(img *image.RGBA, rampLUT *[101][]byte) []byte {
pix := img.Pix
maxCharBytes := 1
for _, char := range rampLUT {
maxCharBytes = max(maxCharBytes, len(char))
}
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"
var ansi256Levels = [6]int{0, 95, 135, 175, 215, 255}
var decimal = func() (t [256]string) {
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
for i, l := range ansi256Levels {
d := v - l
if d < 0 {
d = -d
}
if d < bestDist {
bestDist, best = d, i
}
}
t[v] = uint8(best)
}
return
}()
func quantize256(r, g, b uint8) int {
return 16 + 36*int(ansi256Cube[r]) + 6*int(ansi256Cube[g]) + int(ansi256Cube[b])
}
func appendColor(buf []byte, r, g, b uint8, tier ColorTier) []byte {
if tier == ColorTierTrueColor {
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()
bytesPerCell := 11
if tier == ColorTierTrueColor {
bytesPerCell = 16
}
buf := getOutBuf(bounds.Dx() * bounds.Dy() * bytesPerCell)
var lastR, lastG, lastB uint8
last256 := -1
first := true
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++ {
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
colorChanged := first || r != lastR || g != lastG || bl != lastB
if tier == ColorTier256 {
index := quantize256(r, g, bl)
colorChanged = first || index != last256
last256 = index
}
if colorChanged {
buf = appendColor(buf, r, g, bl, tier)
lastR, lastG, lastB = r, g, bl
first = false
}
buf = append(buf, rampLUT[brightness]...)
o += 4
}
if y < bounds.Max.Y-1 {
buf = append(buf, "\r\n"...)
}
}
buf = append(buf, ansiReset...)
output := bytes.Clone(buf)
putOutBuf(buf)
return output
}
+69
View File
@@ -0,0 +1,69 @@
package render
import (
"path/filepath"
"sync/atomic"
"testing"
"github.com/YuzuZensai/TrollSSH/internal/tsf"
)
func loadBenchSet(b *testing.B) *tsf.FramesContainer {
b.Helper()
matches, _ := filepath.Glob("../../frames/*.tsf")
if len(matches) == 0 {
b.Skip("no .tsf frame set in ../../frames")
}
fc, err := tsf.Load(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 := NewRenderer(0, fc.ColorFrames, Options{
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 := NewRenderer(0, fc.ColorFrames, Options{
BrightnessThreshold: 40,
Charset: "detailed",
}, NewCache(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())
}
}
+194
View File
@@ -0,0 +1,194 @@
package render
import (
"bytes"
"image"
"image/jpeg"
"strings"
"sync"
"testing"
)
func TestResolveCharset(t *testing.T) {
if got := resolveCharset("blocks"); got != " ░▒▓█" {
t.Errorf("blocks preset = %q", got)
}
if got := resolveCharset("XYZ"); got != "XYZ" {
t.Errorf("custom ramp = %q", got)
}
if got := resolveCharset(""); !strings.HasPrefix(got, " .") {
t.Errorf("default = %q", got)
}
}
func TestFrameToAscii(t *testing.T) {
// Below threshold -> first ramp char; full brightness -> last.
opts := Options{BrightnessThreshold: 40, Charset: "standard"}
ramp := []rune(resolveCharset("standard"))
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] {
t.Errorf("dark px = %q, want %q", out[0], ramp[0])
}
if out[1] != ramp[len(ramp)-1] {
t.Errorf("bright px = %q, want %q", out[1], ramp[len(ramp)-1])
}
}
func TestFrameToAsciiInvert(t *testing.T) {
opts := Options{BrightnessThreshold: 40, Charset: "standard", Invert: true}
ramp := []rune(resolveCharset("standard"))
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] {
t.Errorf("inverted bright = %q, want %q", out[0], ramp[0])
}
}
func TestRenderConcurrentSameKey(t *testing.T) {
var jpegBuf bytes.Buffer
src := image.NewRGBA(image.Rect(0, 0, 16, 16))
for i := range src.Pix {
src.Pix[i] = byte(i * 7)
}
if err := jpeg.Encode(&jpegBuf, src, nil); err != nil {
t.Fatal(err)
}
r := NewRenderer(0, [][]byte{jpegBuf.Bytes()}, Options{
BrightnessThreshold: 40,
Charset: "standard",
}, NewCache(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)
}
}
}
func TestRenderCacheEvictsByBytes(t *testing.T) {
key := func(index int) cacheKey { return cacheKey{index: index} }
budget := 3 * entryCost(key(0), bytes.Repeat([]byte("x"), 1000))
c := NewCache(budget)
for i := range 5 {
c.put(key(i), bytes.Repeat([]byte("x"), 1000))
}
if c.size.Load() > budget {
t.Errorf("size %d exceeds budget %d", c.size.Load(), budget)
}
if _, ok := c.get(key(0)); ok {
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 := NewCache(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 := NewCache(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 TestRenderCacheAccountsRetainedCapacity(t *testing.T) {
cache := NewCache(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 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(" .#"), Options{}), 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(" .#"), Options{}), ColorTierTrueColor)
if count := bytes.Count(output, []byte(ansiReset)); count != 1 {
t.Fatalf("reset count = %d, want 1: %q", count, output)
}
}
func TestDetectColorTier(t *testing.T) {
cases := map[string]ColorTier{
"": ColorTierNone,
"dumb": ColorTierNone,
"vt100": ColorTierNone,
"linux": ColorTierNone,
"xterm": ColorTierTrueColor,
"xterm-256color": ColorTier256,
"screen-256color": ColorTier256,
"tmux-256color": ColorTier256,
"xterm-direct": ColorTierTrueColor,
"xterm-kitty": ColorTierTrueColor,
}
for term, want := range cases {
if got := DetectColorTier(term); got != want {
t.Errorf("DetectColorTier(%q) = %d, want %d", term, got, want)
}
}
}
func TestQuantize256(t *testing.T) {
if got := quantize256(0, 0, 0); got != 16 {
t.Errorf("black = %d, want 16", got)
}
if got := quantize256(255, 255, 255); got != 231 {
t.Errorf("white = %d, want 231", got)
}
}
+61
View File
@@ -0,0 +1,61 @@
package sshserver
import (
"crypto/ed25519"
"crypto/rand"
"crypto/rsa"
"encoding/pem"
"fmt"
"os"
"path/filepath"
"golang.org/x/crypto/ssh"
)
func generateAndSave(keyPath, keyType string) error {
fmt.Printf("Generating %s host key...\n", keyType)
var key any
var err error
if keyType == "rsa" {
key, err = rsa.GenerateKey(rand.Reader, 4096)
} else {
_, key, err = ed25519.GenerateKey(rand.Reader)
}
if err != nil {
return err
}
block, err := ssh.MarshalPrivateKey(key, "")
if err != nil {
return err
}
return os.WriteFile(keyPath, pem.EncodeToMemory(block), 0o600)
}
func EnsureHostKeys(configDir string) ([]ssh.Signer, error) {
keys := []struct{ file, keyType string }{
{"id_rsa", "rsa"},
{"id_ed25519", "ed25519"},
}
signers := make([]ssh.Signer, 0, len(keys))
for _, k := range keys {
keyPath := filepath.Join(configDir, k.file)
if _, err := os.Stat(keyPath); os.IsNotExist(err) {
if err := generateAndSave(keyPath, k.keyType); err != nil {
return nil, fmt.Errorf("failed to generate %s host key: %w", k.keyType, err)
}
}
raw, err := os.ReadFile(keyPath)
if err != nil {
return nil, err
}
signer, err := ssh.ParsePrivateKey(raw)
if err != nil {
return nil, fmt.Errorf("failed to parse host key %q: %w", keyPath, err)
}
signers = append(signers, signer)
}
return signers, nil
}
+778
View File
@@ -0,0 +1,778 @@
package sshserver
import (
"encoding/binary"
"errors"
"fmt"
"io"
"math"
"math/rand"
"net"
"strings"
"sync"
"time"
"golang.org/x/crypto/ssh"
"github.com/YuzuZensai/TrollSSH/internal/config"
"github.com/YuzuZensai/TrollSSH/internal/logx"
"github.com/YuzuZensai/TrollSSH/internal/render"
"github.com/YuzuZensai/TrollSSH/internal/tsf"
)
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 {
mu sync.Mutex
counts map[string]int
total int
}
func newConnectionTracker() *ConnectionTracker {
return &ConnectionTracker{counts: make(map[string]int)}
}
func (t *ConnectionTracker) tryAcquire(ip string, maxPerIP, maxTotal int) (int, int, bool) {
t.mu.Lock()
defer t.mu.Unlock()
if t.total >= maxTotal || t.counts[ip] >= maxPerIP {
return t.counts[ip], t.total, false
}
t.counts[ip]++
t.total++
return t.counts[ip], t.total, true
}
func (t *ConnectionTracker) release(ip string) {
t.mu.Lock()
defer t.mu.Unlock()
if _, ok := t.counts[ip]; !ok {
return
}
t.counts[ip]--
if t.total > 0 {
t.total--
}
if t.counts[ip] <= 0 {
delete(t.counts, ip)
}
}
func (t *ConnectionTracker) totalCount() int {
t.mu.Lock()
defer t.mu.Unlock()
return t.total
}
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()
defer t.mu.Unlock()
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 {
data *tsf.FramesContainer
renderer *render.Renderer
}
type Server struct {
config config.Config
sshConfig *ssh.ServerConfig
sets []frameSet
cache *render.Cache
tracker *ConnectionTracker
sessions *SessionTracker
fakeLogin *string
goodbye *string
mu sync.Mutex
listener net.Listener
conns map[net.Conn]struct{}
connWG sync.WaitGroup
closing bool
closeOnce sync.Once
}
type ServerDeps struct {
Config config.Config
HostKeys []ssh.Signer
BannerText *string
FakeLoginText *string
GoodbyeText *string
VideoSets []*tsf.FramesContainer
}
func clampTermSize(cols, rows, maxDimension, maxCells, quantum int) (int, int) {
cols = max(cols, 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))
}
cols = max(1, int(math.Floor(float64(cols)*scale)))
rows = max(1, int(math.Floor(float64(rows)*scale)))
if quantum > 1 {
if cols >= quantum {
cols -= cols % quantum
}
if rows >= quantum {
rows -= rows % quantum
}
}
return cols, rows
}
func New(deps ServerDeps) *Server {
cfg := deps.Config
cache := render.NewCache(int64(cfg.RenderCacheMB) << 20)
sets := make([]frameSet, len(deps.VideoSets))
for i, data := range deps.VideoSets {
sets[i] = frameSet{
data: data,
renderer: render.NewRenderer(i, data.ColorFrames, render.Options{
BrightnessThreshold: cfg.BrightnessThreshold,
Charset: cfg.Charset,
Invert: cfg.Invert,
}, cache),
}
}
sshConfig := &ssh.ServerConfig{
MaxAuthTries: cfg.MaxAuthAttempts,
PasswordCallback: func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) {
ip := hostOnly(conn.RemoteAddr().String())
if cfg.LogCredentials {
logx.Info(fmt.Sprintf(
`Auth attempt from %s method=password user="%s" pass="%s"`,
ip, logx.SanitizeN(conn.User(), 128), logx.SanitizeN(string(password), 128),
))
}
if conn.User() == "" || len(password) == 0 {
return nil, errors.New("password rejected")
}
return nil, nil
},
}
sshConfig.KeyExchanges = []string{
"mlkem768x25519-sha256",
"curve25519-sha256",
"curve25519-sha256@libssh.org",
"ecdh-sha2-nistp256",
"ecdh-sha2-nistp384",
"ecdh-sha2-nistp521",
"diffie-hellman-group14-sha256",
}
if deps.BannerText != nil {
banner := *deps.BannerText
sshConfig.BannerCallback = func(_ ssh.ConnMetadata) string { return banner }
}
for _, key := range deps.HostKeys {
sshConfig.AddHostKey(key)
}
return &Server{
config: cfg,
sshConfig: sshConfig,
sets: sets,
cache: cache,
tracker: newConnectionTracker(),
sessions: newSessionTracker(),
fakeLogin: deps.FakeLoginText,
goodbye: deps.GoodbyeText,
conns: make(map[net.Conn]struct{}),
}
}
func hostOnly(addr string) string {
host, _, err := net.SplitHostPort(addr)
if err != nil {
return addr
}
return host
}
func (s *Server) Listen(host string, port int) error {
listener, err := net.Listen("tcp", net.JoinHostPort(host, fmt.Sprint(port)))
if err != nil {
return err
}
s.mu.Lock()
if s.closing {
s.mu.Unlock()
_ = listener.Close()
return nil
}
s.listener = listener
s.mu.Unlock()
logx.Info(fmt.Sprintf("TrollSSH listening on %s:%d", host, port))
for {
conn, err := listener.Accept()
if err != nil {
if errors.Is(err, net.ErrClosed) {
return nil
}
return err
}
ip := hostOnly(conn.RemoteAddr().String())
activeForIP, total, ok := s.tracker.tryAcquire(ip, s.config.MaxConnections, s.config.MaxTotalConnections)
if !ok {
_ = conn.Close()
logx.Warn("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() {
s.closeOnce.Do(func() {
s.mu.Lock()
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 {
logx.Info(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 {
logx.Warn("Failed to release frame set", set.data.Name, logx.Sanitize(err.Error()))
}
}
})
}
func (s *Server) handleConn(conn net.Conn, ip string, activeForIP, total int) {
defer func() {
s.tracker.release(ip)
s.mu.Lock()
delete(s.conns, conn)
s.mu.Unlock()
s.connWG.Done()
}()
if s.config.HandshakeTimeout > 0 {
_ = conn.SetDeadline(time.Now().Add(s.config.HandshakeTimeout))
}
sshConn, chans, reqs, err := ssh.NewServerConn(conn, s.sshConfig)
if err != nil {
if strings.Contains(err.Error(), "i/o timeout") {
logx.Warn("Handshake timeout for", ip)
} else {
logx.Warn(fmt.Sprintf("Client error from %s:", ip), logx.Sanitize(err.Error()))
}
_ = conn.Close()
return
}
_ = conn.SetDeadline(time.Time{})
logx.Debug("Handshake from", ip)
defer func() { _ = sshConn.Close() }()
setIndex := rand.Intn(len(s.sets))
logx.Info(fmt.Sprintf(
"New connection from %s (ip=%d, total=%d) -> playing %q",
ip, activeForIP, total, s.sets[setIndex].data.Name,
))
go ssh.DiscardRequests(reqs)
var sessionWG sync.WaitGroup
for newChannel := range chans {
if newChannel.ChannelType() != "session" {
_ = newChannel.Reject(ssh.UnknownChannelType, "unknown channel type")
continue
}
if !s.sessions.tryAcquire(sshConn, maxSessionsPerConn, s.config.MaxTotalConnections) {
_ = newChannel.Reject(ssh.ResourceShortage, "session limit reached")
continue
}
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()
logx.Info("Client closed connection from", ip)
}
type termSize struct {
mu sync.Mutex
width int
height int
updated time.Time
}
func (t *termSize) set(w, h, maxDimension, maxCells int, force bool) {
t.mu.Lock()
if !force && time.Since(t.updated) < resizeDebounce {
t.mu.Unlock()
return
}
t.width, t.height = clampTermSize(w, h, maxDimension, maxCells, terminalSizeQuantum)
t.updated = time.Now()
t.mu.Unlock()
}
func (t *termSize) get() (int, int) {
t.mu.Lock()
defer t.mu.Unlock()
return t.width, t.height
}
func parseDims(payload []byte) (cols, rows int, ok bool) {
if len(payload) < 8 {
return 0, 0, false
}
// pty-req prefixes cols/rows with a TERM string; window-change does not.
offset := 0
strLen := binary.BigEndian.Uint32(payload)
if int(strLen)+12 <= len(payload) {
offset = 4 + int(strLen)
}
if len(payload) < offset+8 {
return 0, 0, false
}
cols = int(binary.BigEndian.Uint32(payload[offset:]))
rows = int(binary.BigEndian.Uint32(payload[offset+4:]))
return cols, rows, true
}
// parsePtyTerm extracts the TERM string prefixing a pty-req payload.
func parsePtyTerm(payload []byte) (term string, ok bool) {
if len(payload) < 4 {
return "", false
}
strLen := binary.BigEndian.Uint32(payload)
if int(strLen)+16 > len(payload) {
return "", false
}
return string(payload[4 : 4+strLen]), true
}
func (s *Server) handleSession(
sshConn *ssh.ServerConn,
channel ssh.Channel,
requests <-chan *ssh.Request,
ip string,
initialSetIndex int,
) {
defer func() { _ = channel.Close() }()
size := &termSize{}
size.set(80, 24, s.config.MaxDimension, s.config.MaxTerminalCells, true)
tier := render.ColorTierTrueColor
if s.config.ForceGrayscale {
tier = render.ColorTierNone
}
started := false
var playDone chan struct{}
for req := range requests {
switch req.Type {
case "pty-req":
logx.Debug("Opening pty for session", ip)
if cols, rows, ok := parseDims(req.Payload); ok {
size.set(cols, rows, s.config.MaxDimension, s.config.MaxTerminalCells, true)
}
if term, ok := parsePtyTerm(req.Payload); ok {
tier = render.DetectColorTier(term)
if s.config.ForceGrayscale {
tier = render.ColorTierNone
}
logx.Debug(fmt.Sprintf("Client %s TERM=%q -> color tier %d", ip, logx.SanitizeN(term, 64), tier))
}
_ = req.Reply(true, nil)
case "window-change":
if len(req.Payload) >= 8 {
cols := int(binary.BigEndian.Uint32(req.Payload))
rows := int(binary.BigEndian.Uint32(req.Payload[4:]))
size.set(cols, rows, s.config.MaxDimension, s.config.MaxTerminalCells, false)
}
if req.WantReply {
_ = req.Reply(true, nil)
}
case "exec":
command := ""
if len(req.Payload) >= 4 {
n := binary.BigEndian.Uint32(req.Payload)
if int(n)+4 <= len(req.Payload) {
command = string(req.Payload[4 : 4+n])
}
}
logx.Info(fmt.Sprintf("Client %s attempted exec: %q", ip, logx.SanitizeN(command, 512)))
_ = req.Reply(true, nil)
if !started {
started = true
playDone = make(chan struct{})
playTier := tier
go func(tier render.ColorTier) {
defer close(playDone)
s.playVideo(sshConn, channel, size, ip, initialSetIndex, false, tier)
}(playTier)
}
case "shell":
logx.Debug("Opening shell for session", ip)
_ = req.Reply(true, nil)
if !started {
started = true
playDone = make(chan struct{})
playTier := tier
go func(tier render.ColorTier) {
defer close(playDone)
s.playVideo(sshConn, channel, size, ip, initialSetIndex, false, tier)
}(playTier)
}
default:
if req.WantReply {
_ = req.Reply(false, nil)
}
}
}
_ = channel.Close()
if playDone != nil {
<-playDone
}
}
func (s *Server) pickNextSetIndex(exclude int) int {
if len(s.sets) <= 1 {
return exclude
}
next := exclude
for next == exclude {
next = rand.Intn(len(s.sets))
}
return next
}
func (s *Server) playVideo(
sshConn *ssh.ServerConn,
channel ssh.Channel,
size *termSize,
ip string,
setIndex int,
keepAspectRatio bool,
tier render.ColorTier,
) {
cfg := s.config
current := s.sets[setIndex]
w, h := size.get()
logx.Debug(fmt.Sprintf("Terminal size %dx%d for %s", w, h, ip))
defer func() {
_ = writePartsWithTimeout(sshConn, channel, outputStallTimeout, showCursor)
}()
if s.fakeLogin != nil {
if err := writePartsWithTimeout(
sshConn, channel, outputStallTimeout, clearScreen, *s.fakeLogin,
); err != nil {
return
}
}
done := make(chan struct{})
var doneOnce sync.Once
closeSession := func() {
doneOnce.Do(func() { close(done) })
}
switchCh := make(chan int, 8)
go func() {
buf := make([]byte, 256)
var lastSwitch time.Time
for {
n, err := channel.Read(buf)
if err != nil {
closeSession()
return
}
if !cfg.AllowUserControl {
continue
}
str := string(buf[:n])
delta := 0
if strings.Contains(str, "\x1b[C") || strings.Contains(str, "\x1b[A") {
delta = 1
} else if strings.Contains(str, "\x1b[D") || strings.Contains(str, "\x1b[B") {
delta = -1
}
if delta == 0 {
continue
}
now := time.Now()
if now.Sub(lastSwitch) < cfg.SwitchDebounce {
continue
}
lastSwitch = now
select {
case switchCh <- delta:
default:
}
}
}()
loginTimer := time.NewTimer(cfg.LoginDelay)
select {
case <-loginTimer.C:
case <-done:
if !loginTimer.Stop() {
<-loginTimer.C
}
return
}
if err := writePartsWithTimeout(sshConn, channel, outputStallTimeout, hideCursor); err != nil {
return
}
frameInterval := func() time.Duration {
return time.Duration(float64(time.Second) / current.data.FPS)
}
ticker := time.NewTicker(frameInterval())
defer ticker.Stop()
currentFrame := 0
loopCount := 0
lastW, lastH := 0, 0
for {
select {
case <-done:
return
case delta := <-switchCh:
if len(s.sets) <= 1 {
continue
}
setIndex = (setIndex + delta + len(s.sets)) % len(s.sets)
current = s.sets[setIndex]
currentFrame = 0
lastW, lastH = 0, 0
logx.Debug(fmt.Sprintf("%s switched to %q", ip, current.data.Name))
ticker.Reset(frameInterval())
case <-ticker.C:
w, h := size.get()
ascii, err := current.renderer.Render(currentFrame, w, h, keepAspectRatio, tier)
if err != nil {
logx.Error("Render error for", ip, logx.Sanitize(err.Error()))
_ = sshConn.Close()
return
}
prefix := homeCursor
if w != lastW || h != lastH {
prefix = clearScreen
lastW, lastH = w, h
}
if err := writeFrameWithTimeout(
sshConn, channel, outputStallTimeout, prefix, ascii,
); err != nil {
closeSession()
return
}
currentFrame++
if currentFrame < len(current.data.ColorFrames) {
continue
}
currentFrame = 0
loopCount++
if cfg.MaxLoop > 0 && loopCount >= cfg.MaxLoop {
if err := writePartsWithTimeout(
sshConn, channel, outputStallTimeout, showCursor, clearScreen,
); err != nil {
return
}
if s.goodbye != nil {
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
}
logx.Info("Playback finished, closing session", ip)
_ = channel.Close()
_ = sshConn.Close()
return
}
if cfg.PlaybackMode == config.PlaybackRandom {
setIndex = s.pickNextSetIndex(setIndex)
current = s.sets[setIndex]
logx.Info(fmt.Sprintf(
"Playthrough done for %s, switching to %q", ip, current.data.Name,
))
ticker.Reset(frameInterval())
} else if cfg.MaxLoop > 0 {
logx.Info(fmt.Sprintf(
"Playthrough done for %s, looping %q (%d/%d)",
ip, current.data.Name, loopCount, cfg.MaxLoop,
))
} else {
logx.Info(fmt.Sprintf(
"Playthrough done for %s, looping %q (%d)",
ip, current.data.Name, loopCount,
))
}
}
}
}
+137
View File
@@ -0,0 +1,137 @@
package sshserver
import (
"sync"
"testing"
"golang.org/x/crypto/ssh"
)
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 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 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)
}
}
func TestParseDimsPtyReq(t *testing.T) {
// "xterm" + cols=100 rows=40 + widthpx + heightpx
payload := []byte{
0, 0, 0, 5, 'x', 't', 'e', 'r', 'm',
0, 0, 0, 100,
0, 0, 0, 40,
0, 0, 0, 0,
0, 0, 0, 0,
}
cols, rows, ok := parseDims(payload)
if !ok || cols != 100 || rows != 40 {
t.Errorf("parseDims = %d,%d,%v", cols, rows, ok)
}
term, ok := parsePtyTerm(payload)
if !ok || term != "xterm" {
t.Errorf("parsePtyTerm = %q,%v", term, ok)
}
}
+52
View File
@@ -0,0 +1,52 @@
// .tsf layout, little-endian: "TSFR" | version uint16 | fps float64 |
// count uint32 | count × (colorLen uint32, color JPEG).
package tsf
import "sync"
const (
tsfMagic = "TSFR"
tsfVersion = 1
maxTSFFPS = 240
maxTSFFrameCount = 10_000_000
)
type FramesContainer struct {
ColorFrames [][]byte
FPS float64
Name string
}
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()
}
+119
View File
@@ -0,0 +1,119 @@
package tsf
import (
"bufio"
"encoding/binary"
"fmt"
"math"
"os"
)
func Write(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)
if err != nil {
return err
}
defer func() { _ = f.Close() }()
w := bufio.NewWriterSize(f, 1<<20)
if _, err := w.WriteString(tsfMagic); err != nil {
return err
}
var hdr [14]byte
binary.LittleEndian.PutUint16(hdr[0:], tsfVersion)
binary.LittleEndian.PutUint64(hdr[2:], math.Float64bits(data.FPS))
binary.LittleEndian.PutUint32(hdr[10:], uint32(len(data.ColorFrames)))
if _, err := w.Write(hdr[:]); err != nil {
return err
}
var lenBuf [4]byte
for _, frame := range data.ColorFrames {
binary.LittleEndian.PutUint32(lenBuf[:], uint32(len(frame)))
if _, err := w.Write(lenBuf[:]); err != nil {
return err
}
if _, err := w.Write(frame); err != nil {
return err
}
}
return w.Flush()
}
func Load(filename string) (*FramesContainer, error) {
file, err := readFrameFile(filename)
if err != nil {
return nil, err
}
owned := false
defer func() {
if !owned {
_ = file.Close()
}
}()
raw := file.data
invalid := func() error {
return fmt.Errorf("invalid frames file %q: corrupt .tsf container", filename)
}
if len(raw) < 18 || string(raw[:4]) != tsfMagic {
return nil, invalid()
}
version := binary.LittleEndian.Uint16(raw[4:])
if version != tsfVersion {
return nil, fmt.Errorf("unsupported .tsf version %d in %q", version, filename)
}
fps := math.Float64frombits(binary.LittleEndian.Uint64(raw[6:]))
count := binary.LittleEndian.Uint32(raw[14:])
if math.IsNaN(fps) || math.IsInf(fps, 0) || fps <= 0 || fps > maxTSFFPS {
return nil, fmt.Errorf(
"invalid frames file %q: fps must be finite, greater than 0, and at most %d",
filename, maxTSFFPS,
)
}
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
}
+10
View File
@@ -0,0 +1,10 @@
//go:build !unix
package tsf
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 tsf
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
}
+270
View File
@@ -0,0 +1,270 @@
package tsf
import (
"bytes"
"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 TestTSFRoundTrip(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "f.tsf")
original := &FramesContainer{
ColorFrames: [][]byte{{100, 101, 102}, {110, 120, 130}},
FPS: 29.97,
}
if err := Write(path, original); err != nil {
t.Fatalf("Write: %v", err)
}
fc, err := Load(path)
if err != nil {
t.Fatalf("Load: %v", err)
}
defer func() { _ = fc.Close() }()
if fc.FPS != 29.97 {
t.Errorf("fps = %v", fc.FPS)
}
if len(fc.ColorFrames) != 2 {
t.Fatalf("frames = %d color", len(fc.ColorFrames))
}
if string(fc.ColorFrames[0]) != string([]byte{100, 101, 102}) {
t.Errorf("color frame0 = %v", fc.ColorFrames[0])
}
if string(fc.ColorFrames[1]) != string([]byte{110, 120, 130}) {
t.Errorf("color frame1 = %v", fc.ColorFrames[1])
}
}
func TestTSFInvalid(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "bad.tsf")
if err := os.WriteFile(path, []byte("not a tsf file"), 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
if _, err := Load(path); err == nil {
t.Error("expected error for garbage input")
}
// Valid container but no frames.
if err := Write(path, &FramesContainer{FPS: 30}); err != nil {
t.Fatalf("Write: %v", err)
}
if _, err := Load(path); err == nil {
t.Error("expected error for empty frames")
}
// Valid container but fps <= 0.
rawInvalidFPS := append(tsfHeader(0, 1), 1, 0, 0, 0, 1)
if err := os.WriteFile(path, rawInvalidFPS, 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
if _, err := Load(path); err == nil {
t.Error("expected error for fps<=0")
}
// Truncated payload.
if err := Write(path, &FramesContainer{ColorFrames: [][]byte{{1, 2, 3, 4}}, FPS: 30}); err != nil {
t.Fatalf("Write: %v", err)
}
raw, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile: %v", err)
}
if err := os.WriteFile(path, raw[:len(raw)-2], 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
if _, err := Load(path); err == nil {
t.Error("expected error for truncated file")
}
}
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 := Load(writeRawTSF(t, raw)); err == nil {
t.Errorf("Load accepted fps %v", fps)
}
}
}
func TestTSFRejectsImpossibleCountsAndLengths(t *testing.T) {
if _, err := Load(writeRawTSF(t, tsfHeader(30, math.MaxUint32))); err == nil {
t.Fatal("Load accepted impossible frame count")
}
raw := append(tsfHeader(30, 1), 0xff, 0xff, 0xff, 0xff)
if _, err := Load(writeRawTSF(t, raw)); err == nil {
t.Fatal("Load accepted overflowing frame length")
}
}
func TestTSFCloseReleasesOwnedFrames(t *testing.T) {
path := filepath.Join(t.TempDir(), "frames.tsf")
if err := Write(path, &FramesContainer{FPS: 30, ColorFrames: [][]byte{{1, 2, 3}}}); err != nil {
t.Fatal(err)
}
frames, err := Load(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 := Write(path, &FramesContainer{FPS: math.NaN(), ColorFrames: [][]byte{{1}}})
if err == nil || !strings.Contains(err.Error(), "fps") {
t.Fatalf("Write error = %v", err)
}
if _, err := os.Stat(path); !os.IsNotExist(err) {
t.Fatalf("invalid write created output: %v", err)
}
}
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 := Load(output)
if err != nil {
t.Fatalf("Load: %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)
}
}
+371
View File
@@ -0,0 +1,371 @@
package tsf
import (
"bufio"
"bytes"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"math"
"os"
"os/exec"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/YuzuZensai/TrollSSH/internal/logx"
)
var (
jpegSOI = []byte{0xff, 0xd8}
jpegEOI = []byte{0xff, 0xd9}
)
const (
maxJPEGFrameBytes = 64 << 20
maxFFmpegLogBytes = 64 << 10
)
type jpegFrameSplitter struct {
buffer []byte
scan int
inJPEG bool
}
func (s *jpegFrameSplitter) push(chunk []byte, emit func([]byte) error) (int, error) {
s.buffer = append(s.buffer, chunk...)
emitted := 0
for {
if !s.inJPEG {
start := bytes.Index(s.buffer[s.scan:], jpegSOI)
if start == -1 {
// 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]
}
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 len(s.buffer) > maxJPEGFrameBytes {
return emitted, fmt.Errorf("JPEG frame exceeds %d MiB limit", maxJPEGFrameBytes>>20)
}
s.scan = max(len(jpegSOI), len(s.buffer)-1)
return emitted, nil
}
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.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 = ""
}
}
type ffprobeOutput struct {
Streams []struct {
RFrameRate string `json:"r_frame_rate"`
NbFrames string `json:"nb_frames"`
Duration string `json:"duration"`
} `json:"streams"`
Format struct {
Duration string `json:"duration"`
} `json:"format"`
}
func parseFrameRate(rate string) float64 {
if rate == "" {
return math.NaN()
}
parts := strings.SplitN(rate, "/", 2)
num, err := strconv.ParseFloat(parts[0], 64)
if err != nil {
return math.NaN()
}
if len(parts) == 2 {
den, err := strconv.ParseFloat(parts[1], 64)
if err != nil || den == 0 {
return math.NaN()
}
return num / den
}
return num
}
func extractFrames(path, vf, label string, totalFrames int, emit func([]byte) error) (int, error) {
cmd := exec.Command(
"ffmpeg", "-hide_banner", "-loglevel", "error", "-nostats",
"-i", path,
"-c:v", "mjpeg",
"-q:v", "3",
"-vf", vf,
"-f", "image2pipe",
"pipe:1",
)
stderr := &boundedLog{limit: maxFFmpegLogBytes}
cmd.Stderr = stderr
stdout, err := cmd.StdoutPipe()
if err != nil {
return 0, fmt.Errorf("ffmpeg failed: %w", err)
}
if err := cmd.Start(); err != nil {
return 0, fmt.Errorf("ffmpeg failed: %w", err)
}
splitter := &jpegFrameSplitter{}
frameCount := 0
lastReport := time.Now()
reportProgress := func(force bool) {
if !force && time.Since(lastReport) < 250*time.Millisecond {
return
}
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)
for {
n, readErr := stdout.Read(buf)
if n > 0 {
emitted, splitErr := splitter.push(buf[:n], emit)
frameCount += emitted
if splitErr != nil {
return failStream(fmt.Errorf("ffmpeg stream error: %w", splitErr))
}
if emitted > 0 {
reportProgress(false)
}
}
if readErr == io.EOF {
break
}
if readErr != nil {
return failStream(fmt.Errorf("ffmpeg stream error: %w", readErr))
}
}
if err := cmd.Wait(); err != nil {
if msg := stderr.String(); msg != "" {
return frameCount, fmt.Errorf("ffmpeg failed: %s", msg)
}
return frameCount, fmt.Errorf("ffmpeg failed: %w", err)
}
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()
return frameCount, nil
}
func ProcessVideo(path, output string, maxDimension int) error {
probeCmd := exec.Command(
"ffprobe", "-v", "error",
"-show_streams", "-show_format",
"-of", "json", path,
)
probeOut, err := probeCmd.Output()
if err != nil {
msg := err.Error()
var exitErr *exec.ExitError
if errors.As(err, &exitErr) && len(exitErr.Stderr) > 0 {
msg = strings.TrimSpace(string(exitErr.Stderr))
}
return fmt.Errorf("ffprobe failed: %s", msg)
}
var probe ffprobeOutput
if err := json.Unmarshal(probeOut, &probe); err != nil {
return fmt.Errorf("ffprobe failed: %w", err)
}
if len(probe.Streams) == 0 {
return fmt.Errorf("unable to determine a valid video fps")
}
stream := probe.Streams[0]
fps := parseFrameRate(stream.RFrameRate)
if math.IsNaN(fps) || fps <= 0 {
return fmt.Errorf("unable to determine a valid video fps")
}
totalFrames := 0
if n, err := strconv.Atoi(stream.NbFrames); err == nil {
totalFrames = n
} else {
durStr := probe.Format.Duration
if durStr == "" {
durStr = stream.Duration
}
if d, err := strconv.ParseFloat(durStr, 64); err == nil {
totalFrames = int(math.Round(d * fps))
}
}
scaleFilter := fmt.Sprintf(
"scale=w=%d:h=%d:force_original_aspect_ratio=decrease",
maxDimension, maxDimension,
)
outputFile, err := newStreamingTSF(output, fps)
if err != nil {
return err
}
defer outputFile.abort()
frameCount, err := extractFrames(path, scaleFilter, "color", totalFrames, outputFile.addFrame)
if err != nil {
return err
}
if err := outputFile.commit(output); err != nil {
return err
}
logx.Info(fmt.Sprintf("Saved %d frames to %s", frameCount, output))
return nil
}