Files

269 lines
6.1 KiB
Go
Raw Permalink Normal View History

2026-07-13 17:47:00 +07:00
package main
import (
"fmt"
"os"
"os/signal"
"path/filepath"
"runtime"
2026-07-14 05:49:10 +07:00
"runtime/debug"
2026-07-13 17:47:00 +07:00
"sort"
2026-07-14 05:49:10 +07:00
"strconv"
2026-07-13 17:47:00 +07:00
"strings"
"sync"
"syscall"
"time"
"github.com/joho/godotenv"
2026-07-16 23:29:14 +07:00
"github.com/YuzuZensai/TrollSSH/internal/config"
"github.com/YuzuZensai/TrollSSH/internal/logx"
"github.com/YuzuZensai/TrollSSH/internal/sshserver"
"github.com/YuzuZensai/TrollSSH/internal/tsf"
2026-07-13 17:47:00 +07:00
)
type cliArgs struct {
2026-07-14 05:49:10 +07:00
generate bool
video string
resolution int
2026-07-13 17:47:00 +07:00
}
func parseArgs(argv []string) cliArgs {
2026-07-14 05:49:10 +07:00
args := cliArgs{resolution: 512}
2026-07-13 17:47:00 +07:00
for i := 0; i < len(argv); i++ {
switch argv[i] {
case "--generate", "-g":
args.generate = true
case "--video", "-v":
if i+1 < len(argv) {
i++
args.video = argv[i]
}
2026-07-14 05:49:10 +07:00
case "--resolution", "-r":
if i+1 < len(argv) {
i++
if n, err := strconv.Atoi(argv[i]); err == nil {
args.resolution = max(n, 16)
}
}
2026-07-13 17:47:00 +07:00
}
}
return args
}
func fail(message string) {
2026-07-16 23:29:14 +07:00
logx.Error(message)
2026-07-13 17:47:00 +07:00
os.Exit(1)
}
func resolveVideoPath(explicitPath string) string {
abs, err := filepath.Abs(explicitPath)
if err != nil {
return ""
}
if info, err := os.Stat(abs); err == nil && !info.IsDir() {
return abs
}
return ""
}
2026-07-14 05:49:10 +07:00
func generateFrames(framesDir, videoArg string, resolution int) {
2026-07-13 17:47:00 +07:00
if videoArg == "" {
fail("No source video given. Pass --video <path>.")
}
videoPath := resolveVideoPath(videoArg)
if videoPath == "" {
fail(fmt.Sprintf("Source video %q does not exist or is not a file.", videoArg))
}
if err := os.MkdirAll(framesDir, 0o755); err != nil {
fail(fmt.Sprintf("Failed to create frames directory %q: %s", framesDir, err.Error()))
}
2026-07-13 17:47:00 +07:00
base := strings.TrimSuffix(filepath.Base(videoPath), filepath.Ext(videoPath))
output := filepath.Join(framesDir, base+".tsf")
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf("Generating frames from %q -> %s", videoPath, output))
if err := tsf.ProcessVideo(videoPath, output, resolution); err != nil {
2026-07-13 17:47:00 +07:00
fail(fmt.Sprintf("Failed to generate frames from %q: %s", videoPath, err.Error()))
}
}
const frameDataWarnBytes = 2 << 30
2026-07-16 23:29:14 +07:00
func loadAllFrames(framesDir string) []*tsf.FramesContainer {
2026-07-13 17:47:00 +07:00
entries, err := os.ReadDir(framesDir)
var files []string
if err == nil {
for _, e := range entries {
if strings.HasSuffix(strings.ToLower(e.Name()), ".tsf") {
files = append(files, e.Name())
}
}
sort.Strings(files)
}
if len(files) == 0 {
fail(fmt.Sprintf(
"No frame sets found in %q. Generate one first with: trollssh --generate --video <path>",
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 {
2026-07-16 23:29:14 +07:00
logx.Warn(fmt.Sprintf(
"Frame data is %.1f MB of mapped memory; make sure the container memory limit leaves headroom",
float64(totalBytes)/(1<<20),
))
} else {
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf("Frame data: %.1f MB", float64(totalBytes)/(1<<20)))
}
2026-07-13 17:47:00 +07:00
concurrency := min(len(files), max(1, min(runtime.NumCPU(), 4)))
2026-07-16 23:29:14 +07:00
results := make([]*tsf.FramesContainer, len(files))
2026-07-13 17:47:00 +07:00
errs := make([]error, len(files))
var next int
var nextMu sync.Mutex
var wg sync.WaitGroup
worker := func() {
defer wg.Done()
for {
nextMu.Lock()
i := next
next++
nextMu.Unlock()
if i >= len(files) {
return
}
file := files[i]
filePath := filepath.Join(framesDir, file)
info, err := os.Stat(filePath)
if err != nil {
errs[i] = err
return
}
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf("Loading %s (%.1f MB)...", file, float64(info.Size())/1024/1024))
data, err := tsf.Load(filePath)
2026-07-13 17:47:00 +07:00
if err != nil {
errs[i] = err
return
}
data.Name = file
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf(" %s: %d frames @ %gfps", file, len(data.ColorFrames), data.FPS))
2026-07-13 17:47:00 +07:00
results[i] = data
}
}
wg.Add(concurrency)
for range concurrency {
go worker()
}
wg.Wait()
for _, err := range errs {
if err != nil {
fail(err.Error())
}
}
return results
}
2026-07-14 05:49:10 +07:00
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)
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf("Memory limit set to %d MB (90%% of cgroup limit)", limit>>20))
2026-07-14 05:49:10 +07:00
return
}
}
2026-07-13 17:47:00 +07:00
func main() {
_ = godotenv.Load()
2026-07-16 23:29:14 +07:00
logx.SetThreshold(logx.ResolveThreshold())
2026-07-14 05:49:10 +07:00
applyMemoryLimit()
2026-07-13 17:47:00 +07:00
2026-07-16 23:29:14 +07:00
cfg := config.Load()
2026-07-13 17:47:00 +07:00
args := parseArgs(os.Args[1:])
cwd, err := os.Getwd()
if err != nil {
fail(err.Error())
}
dataDir := filepath.Join(cwd, "data")
framesDir := filepath.Join(cwd, "frames")
if args.generate {
2026-07-14 05:49:10 +07:00
generateFrames(framesDir, args.video, args.resolution)
2026-07-13 17:47:00 +07:00
return
}
if err := os.MkdirAll(dataDir, 0o755); err != nil {
fail(fmt.Sprintf("Failed to create data directory %q: %s", dataDir, err.Error()))
}
2026-07-13 17:47:00 +07:00
var bannerText, fakeLoginText, goodbyeText *string
2026-07-16 23:29:14 +07:00
if text, ok := config.LoadOptionalTextFile(filepath.Join(dataDir, "banner.txt")); ok {
2026-07-13 17:47:00 +07:00
bannerText = &text
}
2026-07-16 23:29:14 +07:00
if text, ok := config.LoadOptionalTextFile(filepath.Join(dataDir, "fakelogin.txt")); ok {
2026-07-13 17:47:00 +07:00
fakeLoginText = &text
}
2026-07-16 23:29:14 +07:00
if text, ok := config.LoadOptionalTextFile(filepath.Join(dataDir, "goodbye.txt")); ok {
2026-07-13 17:47:00 +07:00
goodbyeText = &text
}
2026-07-16 23:29:14 +07:00
hostKeys, err := sshserver.EnsureHostKeys(dataDir)
2026-07-13 17:47:00 +07:00
if err != nil {
fail(err.Error())
}
videoSets := loadAllFrames(framesDir)
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf("Loaded %d frame set(s)", len(videoSets)))
2026-07-13 17:47:00 +07:00
2026-07-16 23:29:14 +07:00
server := sshserver.New(sshserver.ServerDeps{
Config: cfg,
2026-07-13 17:47:00 +07:00
HostKeys: hostKeys,
BannerText: bannerText,
FakeLoginText: fakeLoginText,
GoodbyeText: goodbyeText,
VideoSets: videoSets,
})
defer server.Close()
2026-07-13 17:47:00 +07:00
sigCh := make(chan os.Signal, 1)
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
go func() {
sig := <-sigCh
2026-07-16 23:29:14 +07:00
logx.Info(fmt.Sprintf("Received %s, shutting down...", sig))
forceExit := time.AfterFunc(5*time.Second, func() { os.Exit(0) })
2026-07-13 17:47:00 +07:00
server.Close()
forceExit.Stop()
2026-07-13 17:47:00 +07:00
}()
2026-07-16 23:29:14 +07:00
if err := server.Listen(cfg.Host, cfg.Port); err != nil {
logx.Error("Server error:", err.Error())
2026-07-13 17:47:00 +07:00
os.Exit(1)
}
}