Initial commit: Null DRM Official
Capture, decrypt, and restream toolkit with compiled-in app modules (RTE, TG4, BBC), on-device MITM proxy, streamd control plane, and www. BBC module.yaml is published (clear streams); other module values stay local.
This commit is contained in:
commit
2fa8f2435f
121 changed files with 17802 additions and 0 deletions
118
apps/agent/internal/client/client.go
Normal file
118
apps/agent/internal/client/client.go
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
package client
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
BaseURL string
|
||||
Token string
|
||||
HTTP *http.Client
|
||||
UserAgent string
|
||||
}
|
||||
|
||||
func New(baseURL, token string) *Client {
|
||||
return &Client{
|
||||
BaseURL: strings.TrimRight(baseURL, "/"),
|
||||
Token: token,
|
||||
HTTP: &http.Client{Timeout: 30 * time.Second},
|
||||
UserAgent: "drm-agent/1.0",
|
||||
}
|
||||
}
|
||||
|
||||
type Stream struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Title string `json:"title"`
|
||||
App string `json:"app"`
|
||||
Channel string `json:"channel"`
|
||||
MPD string `json:"mpd"`
|
||||
Key string `json:"key"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Health string `json:"health"`
|
||||
LastError string `json:"last_error"`
|
||||
PlaylistAgeS float64 `json:"playlist_age_s"`
|
||||
ClaimedBy string `json:"claimed_by"`
|
||||
}
|
||||
|
||||
type Credentials struct {
|
||||
MPD string `json:"mpd"`
|
||||
Key string `json:"key"`
|
||||
PSSH string `json:"pssh,omitempty"`
|
||||
Auth string `json:"auth,omitempty"`
|
||||
PID string `json:"pid,omitempty"`
|
||||
}
|
||||
|
||||
func (c *Client) Health() error {
|
||||
_, err := c.do(http.MethodGet, "/api/health", nil)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) ListStreams() ([]Stream, error) {
|
||||
body, err := c.do(http.MethodGet, "/api/streams", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var wrap struct {
|
||||
Streams []Stream `json:"streams"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &wrap); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return wrap.Streams, nil
|
||||
}
|
||||
|
||||
func (c *Client) Claim(id int64, agentID string, ttlSec int) error {
|
||||
payload := map[string]any{"agent_id": agentID, "ttl_sec": ttlSec}
|
||||
_, err := c.do(http.MethodPost, fmt.Sprintf("/api/streams/%d/claim", id), payload)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) ReleaseClaim(id int64, agentID string) error {
|
||||
payload := map[string]any{"agent_id": agentID}
|
||||
_, err := c.do(http.MethodPost, fmt.Sprintf("/api/streams/%d/claim/release", id), payload)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) PostCredentials(id int64, cred Credentials) error {
|
||||
_, err := c.do(http.MethodPost, fmt.Sprintf("/api/streams/%d/credentials", id), cred)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) do(method, path string, payload any) ([]byte, error) {
|
||||
var rdr io.Reader
|
||||
if payload != nil {
|
||||
b, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rdr = bytes.NewReader(b)
|
||||
}
|
||||
req, err := http.NewRequest(method, c.BaseURL+path, rdr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if payload != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
req.Header.Set("User-Agent", c.UserAgent)
|
||||
if c.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+c.Token)
|
||||
}
|
||||
resp, err := c.HTTP.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("%s %s: %s — %s", method, path, resp.Status, strings.TrimSpace(string(body)))
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
104
apps/agent/internal/config/config.go
Normal file
104
apps/agent/internal/config/config.go
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
|
||||
"drmdecryption/repo"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
StreamdURL string `yaml:"streamd_url"`
|
||||
Token string `yaml:"token"`
|
||||
PollIntervalSec int `yaml:"poll_interval_sec"`
|
||||
DataDir string `yaml:"data_dir"`
|
||||
Workers int `yaml:"workers"`
|
||||
MaxAttempts int `yaml:"max_attempts"`
|
||||
CaptureWaitSec int `yaml:"capture_wait_sec"`
|
||||
Devices []string `yaml:"devices"`
|
||||
Python string `yaml:"python"`
|
||||
WVD string `yaml:"wvd"` // path to .wvd device file (required for key fetch)
|
||||
UserAgent string `yaml:"user_agent"`
|
||||
AgentID string `yaml:"agent_id"`
|
||||
}
|
||||
|
||||
func (c Config) PollInterval() time.Duration {
|
||||
if c.PollIntervalSec <= 0 {
|
||||
return 30 * time.Second
|
||||
}
|
||||
return time.Duration(c.PollIntervalSec) * time.Second
|
||||
}
|
||||
|
||||
func (c Config) CaptureWait() time.Duration {
|
||||
if c.CaptureWaitSec <= 0 {
|
||||
return 180 * time.Second
|
||||
}
|
||||
return time.Duration(c.CaptureWaitSec) * time.Second
|
||||
}
|
||||
|
||||
func Load(path string) (Config, error) {
|
||||
root := repo.Root()
|
||||
if path == "" {
|
||||
path = filepath.Join(root, "apps", "agent", "agent.yaml")
|
||||
}
|
||||
var cfg Config
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
cfg = defaults(root)
|
||||
return cfg, nil
|
||||
}
|
||||
return cfg, err
|
||||
}
|
||||
if err := yaml.Unmarshal(raw, &cfg); err != nil {
|
||||
return cfg, err
|
||||
}
|
||||
cfg = applyDefaults(cfg, root)
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func defaults(root string) Config {
|
||||
return applyDefaults(Config{}, root)
|
||||
}
|
||||
|
||||
func applyDefaults(cfg Config, root string) Config {
|
||||
if cfg.StreamdURL == "" {
|
||||
cfg.StreamdURL = "http://127.0.0.1:8083"
|
||||
}
|
||||
if cfg.Token == "" {
|
||||
cfg.Token = os.Getenv("STREAMD_TOKEN")
|
||||
}
|
||||
if cfg.PollIntervalSec <= 0 {
|
||||
cfg.PollIntervalSec = 30
|
||||
}
|
||||
if cfg.DataDir == "" {
|
||||
cfg.DataDir = filepath.Join(root, ".cache", "agent")
|
||||
} else if !filepath.IsAbs(cfg.DataDir) {
|
||||
cfg.DataDir = filepath.Join(root, cfg.DataDir)
|
||||
}
|
||||
if cfg.Workers <= 0 {
|
||||
cfg.Workers = 1
|
||||
}
|
||||
if cfg.MaxAttempts <= 0 {
|
||||
cfg.MaxAttempts = 5
|
||||
}
|
||||
if cfg.CaptureWaitSec <= 0 {
|
||||
cfg.CaptureWaitSec = 180
|
||||
}
|
||||
if cfg.AgentID == "" {
|
||||
// Stable across restarts so the same agent can reclaim after a crash.
|
||||
// Override in apps/agent/agent.yaml when running multiple agents.
|
||||
host, _ := os.Hostname()
|
||||
if host == "" {
|
||||
host = "agent"
|
||||
}
|
||||
cfg.AgentID = host
|
||||
}
|
||||
if cfg.WVD != "" && !filepath.IsAbs(cfg.WVD) {
|
||||
cfg.WVD = filepath.Join(root, cfg.WVD)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
118
apps/agent/internal/device/pool.go
Normal file
118
apps/agent/internal/device/pool.go
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
package device
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"drmdecryption/adb"
|
||||
)
|
||||
|
||||
// Pool tracks free/busy ADB devices. Never installs apps.
|
||||
type Pool struct {
|
||||
mu sync.Mutex
|
||||
adbBin string
|
||||
allow map[string]bool // empty allow = all serials
|
||||
busy map[string]string // serial → job label
|
||||
}
|
||||
|
||||
func New(adbBin string, allowlist []string) *Pool {
|
||||
p := &Pool{
|
||||
adbBin: adbBin,
|
||||
allow: map[string]bool{},
|
||||
busy: map[string]string{},
|
||||
}
|
||||
for _, s := range allowlist {
|
||||
if s != "" {
|
||||
p.allow[s] = true
|
||||
}
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
type Info struct {
|
||||
Serial string `json:"serial"`
|
||||
State string `json:"state"`
|
||||
Model string `json:"model"`
|
||||
Busy bool `json:"busy"`
|
||||
BusyFor string `json:"busy_for,omitempty"`
|
||||
Allowed bool `json:"allowed"`
|
||||
}
|
||||
|
||||
func (p *Pool) List() ([]Info, error) {
|
||||
devs, err := adb.ListDevices(p.adbBin)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
out := make([]Info, 0, len(devs))
|
||||
for _, d := range devs {
|
||||
allowed := len(p.allow) == 0 || p.allow[d.Serial]
|
||||
busyFor, busy := p.busy[d.Serial]
|
||||
out = append(out, Info{
|
||||
Serial: d.Serial,
|
||||
State: d.State,
|
||||
Model: d.Model,
|
||||
Busy: busy,
|
||||
BusyFor: busyFor,
|
||||
Allowed: allowed,
|
||||
})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Acquire finds a free device in state "device" that already has pkg installed.
|
||||
// Returns ErrWait if none available (do not burn job attempts).
|
||||
func (p *Pool) Acquire(pkg, jobLabel string) (*adb.Client, string, error) {
|
||||
devs, err := adb.ListDevices(p.adbBin)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
base := adb.New()
|
||||
if p.adbBin != "" {
|
||||
base.Bin = p.adbBin
|
||||
}
|
||||
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
var lastMiss string
|
||||
for _, d := range devs {
|
||||
if d.State != "device" {
|
||||
continue
|
||||
}
|
||||
if len(p.allow) > 0 && !p.allow[d.Serial] {
|
||||
continue
|
||||
}
|
||||
if _, busy := p.busy[d.Serial]; busy {
|
||||
continue
|
||||
}
|
||||
c := base.WithSerial(d.Serial)
|
||||
if pkg != "" && !c.PackageInstalled(pkg) {
|
||||
lastMiss = fmt.Sprintf("%s missing package %s", d.Serial, pkg)
|
||||
continue
|
||||
}
|
||||
p.busy[d.Serial] = jobLabel
|
||||
return c, d.Serial, nil
|
||||
}
|
||||
if lastMiss != "" {
|
||||
return nil, "", &WaitError{Msg: lastMiss}
|
||||
}
|
||||
return nil, "", &WaitError{Msg: "no free adb device"}
|
||||
}
|
||||
|
||||
func (p *Pool) Release(serial string) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
delete(p.busy, serial)
|
||||
}
|
||||
|
||||
// WaitError means the job should stay queued without consuming an attempt.
|
||||
type WaitError struct{ Msg string }
|
||||
|
||||
func (e *WaitError) Error() string { return e.Msg }
|
||||
|
||||
func IsWait(err error) bool {
|
||||
_, ok := err.(*WaitError)
|
||||
return ok
|
||||
}
|
||||
111
apps/agent/internal/poller/poller.go
Normal file
111
apps/agent/internal/poller/poller.go
Normal file
|
|
@ -0,0 +1,111 @@
|
|||
package poller
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"drmdecryption/apps/agent/internal/client"
|
||||
"drmdecryption/apps/agent/internal/queue"
|
||||
)
|
||||
|
||||
type Poller struct {
|
||||
API *client.Client
|
||||
Queue *queue.Store
|
||||
MaxAttempts int
|
||||
Interval time.Duration
|
||||
Log *log.Logger
|
||||
}
|
||||
|
||||
func (p *Poller) Once() error {
|
||||
streams, err := p.API.ListStreams()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, st := range streams {
|
||||
if !st.Enabled {
|
||||
continue
|
||||
}
|
||||
reason := needsRefresh(st)
|
||||
if reason == "" {
|
||||
continue
|
||||
}
|
||||
j, added, err := p.Queue.EnqueueIfIdle(st.ID, st.Name, st.App, st.Channel, reason, p.MaxAttempts)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if added {
|
||||
p.logf("enqueued job=%d stream=%s reason=%s", j.ID, st.Name, reason)
|
||||
continue
|
||||
}
|
||||
if active, ok, err := p.Queue.ActiveJobForStream(st.ID); err == nil && ok {
|
||||
extra := ""
|
||||
if active.Status == queue.StatusQueued && !active.NotBefore.IsZero() && active.NotBefore.After(time.Now().UTC()) {
|
||||
extra = fmt.Sprintf(" backoff_until=%s", active.NotBefore.Format(time.RFC3339))
|
||||
}
|
||||
p.logf("stream %s needs %s but job=%d already %s (attempts %d/%d)%s",
|
||||
st.Name, reason, active.ID, active.Status, active.Attempts, active.MaxAttempts, extra)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Poller) Loop(stop <-chan struct{}) {
|
||||
t := time.NewTicker(p.Interval)
|
||||
defer t.Stop()
|
||||
_ = p.Once()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
case <-t.C:
|
||||
if err := p.Once(); err != nil {
|
||||
p.logf("poll error: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func needsRefresh(st client.Stream) string {
|
||||
if strings.TrimSpace(st.MPD) == "" || strings.TrimSpace(st.Key) == "" {
|
||||
return "missing_credentials"
|
||||
}
|
||||
if !validKeyPair(st.Key) {
|
||||
return "invalid_credentials"
|
||||
}
|
||||
h := strings.ToLower(strings.TrimSpace(st.Health))
|
||||
switch h {
|
||||
case "down", "error", "failed":
|
||||
return "down"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// validKeyPair accepts KID:KEY as 32-hex:32-hex (colons/dashes in KID ignored).
|
||||
func validKeyPair(key string) bool {
|
||||
parts := strings.SplitN(strings.TrimSpace(key), ":", 2)
|
||||
if len(parts) != 2 {
|
||||
return false
|
||||
}
|
||||
hexOnly := func(s string) string {
|
||||
var b strings.Builder
|
||||
for _, r := range s {
|
||||
switch {
|
||||
case r >= '0' && r <= '9', r >= 'a' && r <= 'f', r >= 'A' && r <= 'F':
|
||||
b.WriteRune(r)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
kid, k := hexOnly(parts[0]), hexOnly(parts[1])
|
||||
return len(kid) == 32 && len(k) == 32
|
||||
}
|
||||
|
||||
func (p *Poller) logf(format string, args ...any) {
|
||||
if p.Log != nil {
|
||||
p.Log.Printf(format, args...)
|
||||
return
|
||||
}
|
||||
log.Printf(format, args...)
|
||||
}
|
||||
334
apps/agent/internal/queue/queue.go
Normal file
334
apps/agent/internal/queue/queue.go
Normal file
|
|
@ -0,0 +1,334 @@
|
|||
package queue
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
type Status string
|
||||
|
||||
const (
|
||||
StatusQueued Status = "queued"
|
||||
StatusRunning Status = "running"
|
||||
StatusOK Status = "ok"
|
||||
StatusFailed Status = "failed"
|
||||
StatusCancel Status = "cancelled"
|
||||
)
|
||||
|
||||
type Job struct {
|
||||
ID int64
|
||||
StreamID int64
|
||||
StreamName string
|
||||
App string
|
||||
Channel string
|
||||
Reason string
|
||||
Status Status
|
||||
Priority int
|
||||
Attempts int
|
||||
MaxAttempts int
|
||||
DeviceSerial string
|
||||
LastError string
|
||||
NotBefore time.Time
|
||||
CreatedAt time.Time
|
||||
StartedAt time.Time
|
||||
FinishedAt time.Time
|
||||
}
|
||||
|
||||
type Store struct {
|
||||
DB *sql.DB
|
||||
}
|
||||
|
||||
func Open(dataDir string) (*Store, error) {
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
path := filepath.Join(dataDir, "agent.db")
|
||||
dsn := fmt.Sprintf("file:%s?_pragma=busy_timeout(5000)&_pragma=foreign_keys(1)", filepath.ToSlash(path))
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
s := &Store{DB: db}
|
||||
if err := s.migrate(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (s *Store) Close() error { return s.DB.Close() }
|
||||
|
||||
func (s *Store) migrate() error {
|
||||
_, err := s.DB.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS jobs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
stream_id INTEGER NOT NULL,
|
||||
stream_name TEXT NOT NULL DEFAULT '',
|
||||
app TEXT NOT NULL DEFAULT '',
|
||||
channel TEXT NOT NULL DEFAULT '',
|
||||
reason TEXT NOT NULL DEFAULT '',
|
||||
status TEXT NOT NULL DEFAULT 'queued',
|
||||
priority INTEGER NOT NULL DEFAULT 100,
|
||||
attempts INTEGER NOT NULL DEFAULT 0,
|
||||
max_attempts INTEGER NOT NULL DEFAULT 5,
|
||||
device_serial TEXT NOT NULL DEFAULT '',
|
||||
last_error TEXT NOT NULL DEFAULT '',
|
||||
not_before TEXT NOT NULL DEFAULT '',
|
||||
created_at TEXT NOT NULL,
|
||||
started_at TEXT NOT NULL DEFAULT '',
|
||||
finished_at TEXT NOT NULL DEFAULT ''
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_jobs_status ON jobs(status, priority, id);
|
||||
CREATE INDEX IF NOT EXISTS idx_jobs_stream ON jobs(stream_id, status);
|
||||
`)
|
||||
return err
|
||||
}
|
||||
|
||||
func now() string { return time.Now().UTC().Format(time.RFC3339) }
|
||||
|
||||
func parseTime(s string) time.Time {
|
||||
if s == "" {
|
||||
return time.Time{}
|
||||
}
|
||||
t, _ := time.Parse(time.RFC3339, s)
|
||||
return t
|
||||
}
|
||||
|
||||
// EnqueueIfIdle inserts a job unless one is already queued/running for the stream.
|
||||
func (s *Store) EnqueueIfIdle(streamID int64, name, app, channel, reason string, maxAttempts int) (Job, bool, error) {
|
||||
var existing int
|
||||
err := s.DB.QueryRow(`
|
||||
SELECT COUNT(1) FROM jobs WHERE stream_id=? AND status IN ('queued','running')`, streamID).Scan(&existing)
|
||||
if err != nil {
|
||||
return Job{}, false, err
|
||||
}
|
||||
if existing > 0 {
|
||||
return Job{}, false, nil
|
||||
}
|
||||
if maxAttempts <= 0 {
|
||||
maxAttempts = 5
|
||||
}
|
||||
res, err := s.DB.Exec(`
|
||||
INSERT INTO jobs(stream_id,stream_name,app,channel,reason,status,priority,attempts,max_attempts,created_at,not_before)
|
||||
VALUES(?,?,?,?,?,'queued',100,0,?,?,?)`,
|
||||
streamID, name, app, channel, reason, maxAttempts, now(), now())
|
||||
if err != nil {
|
||||
return Job{}, false, err
|
||||
}
|
||||
id, _ := res.LastInsertId()
|
||||
j, err := s.Get(id)
|
||||
return j, true, err
|
||||
}
|
||||
|
||||
func (s *Store) Get(id int64) (Job, error) {
|
||||
row := s.DB.QueryRow(`
|
||||
SELECT id,stream_id,stream_name,app,channel,reason,status,priority,attempts,max_attempts,
|
||||
device_serial,last_error,not_before,created_at,started_at,finished_at
|
||||
FROM jobs WHERE id=?`, id)
|
||||
return scanJob(row)
|
||||
}
|
||||
|
||||
type scannable interface {
|
||||
Scan(dest ...any) error
|
||||
}
|
||||
|
||||
func scanJob(row scannable) (Job, error) {
|
||||
var j Job
|
||||
var status, nb, ca, sa, fa string
|
||||
err := row.Scan(
|
||||
&j.ID, &j.StreamID, &j.StreamName, &j.App, &j.Channel, &j.Reason, &status,
|
||||
&j.Priority, &j.Attempts, &j.MaxAttempts, &j.DeviceSerial, &j.LastError,
|
||||
&nb, &ca, &sa, &fa,
|
||||
)
|
||||
if err != nil {
|
||||
return j, err
|
||||
}
|
||||
j.Status = Status(status)
|
||||
j.NotBefore = parseTime(nb)
|
||||
j.CreatedAt = parseTime(ca)
|
||||
j.StartedAt = parseTime(sa)
|
||||
j.FinishedAt = parseTime(fa)
|
||||
return j, nil
|
||||
}
|
||||
|
||||
// ClaimNext marks the next ready queued job as running. Returns false if none.
|
||||
func (s *Store) ClaimNext() (Job, bool, error) {
|
||||
tx, err := s.DB.Begin()
|
||||
if err != nil {
|
||||
return Job{}, false, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
|
||||
nowStr := now()
|
||||
row := tx.QueryRow(`
|
||||
SELECT id FROM jobs
|
||||
WHERE status='queued' AND (not_before='' OR not_before<=?)
|
||||
ORDER BY priority ASC, id ASC
|
||||
LIMIT 1`, nowStr)
|
||||
var id int64
|
||||
if err := row.Scan(&id); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return Job{}, false, nil
|
||||
}
|
||||
return Job{}, false, err
|
||||
}
|
||||
_, err = tx.Exec(`UPDATE jobs SET status='running', started_at=?, last_error='' WHERE id=? AND status='queued'`, nowStr, id)
|
||||
if err != nil {
|
||||
return Job{}, false, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return Job{}, false, err
|
||||
}
|
||||
j, err := s.Get(id)
|
||||
return j, true, err
|
||||
}
|
||||
|
||||
func (s *Store) SetDevice(id int64, serial string) error {
|
||||
_, err := s.DB.Exec(`UPDATE jobs SET device_serial=? WHERE id=?`, serial, id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) MarkOK(id int64) error {
|
||||
_, err := s.DB.Exec(`UPDATE jobs SET status='ok', finished_at=?, last_error='' WHERE id=?`, now(), id)
|
||||
return err
|
||||
}
|
||||
|
||||
// MarkFailed increments attempts. Re-queues with backoff unless maxed out.
|
||||
// waitingDevice=true keeps status queued without burning an attempt.
|
||||
func (s *Store) MarkFailed(id int64, msg string, waitingDevice bool) error {
|
||||
j, err := s.Get(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if waitingDevice {
|
||||
_, err = s.DB.Exec(`
|
||||
UPDATE jobs SET status='queued', device_serial='', last_error=?, started_at='' WHERE id=?`,
|
||||
msg, id)
|
||||
return err
|
||||
}
|
||||
attempts := j.Attempts + 1
|
||||
if attempts >= j.MaxAttempts {
|
||||
_, err = s.DB.Exec(`
|
||||
UPDATE jobs SET status='failed', attempts=?, last_error=?, finished_at=? WHERE id=?`,
|
||||
attempts, msg, now(), id)
|
||||
return err
|
||||
}
|
||||
backoff := time.Duration(1<<uint(min(attempts, 4))) * time.Minute
|
||||
if backoff > 15*time.Minute {
|
||||
backoff = 15 * time.Minute
|
||||
}
|
||||
nb := time.Now().UTC().Add(backoff).Format(time.RFC3339)
|
||||
_, err = s.DB.Exec(`
|
||||
UPDATE jobs SET status='queued', attempts=?, last_error=?, not_before=?, device_serial='', started_at='' WHERE id=?`,
|
||||
attempts, msg, nb, id)
|
||||
return err
|
||||
}
|
||||
|
||||
func min(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func (s *Store) Cancel(id int64) error {
|
||||
_, err := s.DB.Exec(`
|
||||
UPDATE jobs SET status='cancelled', finished_at=? WHERE id=? AND status IN ('queued','running')`, now(), id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) ListRecent(limit int) ([]Job, error) {
|
||||
if limit <= 0 {
|
||||
limit = 30
|
||||
}
|
||||
rows, err := s.DB.Query(`
|
||||
SELECT id,stream_id,stream_name,app,channel,reason,status,priority,attempts,max_attempts,
|
||||
device_serial,last_error,not_before,created_at,started_at,finished_at
|
||||
FROM jobs ORDER BY id DESC LIMIT ?`, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []Job
|
||||
for rows.Next() {
|
||||
j, err := scanJob(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, j)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (s *Store) ActiveForStream(streamID int64) (bool, error) {
|
||||
var n int
|
||||
err := s.DB.QueryRow(`SELECT COUNT(1) FROM jobs WHERE stream_id=? AND status IN ('queued','running')`, streamID).Scan(&n)
|
||||
return n > 0, err
|
||||
}
|
||||
|
||||
// RecoverStaleRunning re-queues jobs left in status=running after a crash/kill
|
||||
// so a down stream is not blocked forever by EnqueueIfIdle.
|
||||
func (s *Store) RecoverStaleRunning() (int, error) {
|
||||
nowStr := now()
|
||||
res, err := s.DB.Exec(`
|
||||
UPDATE jobs SET status='queued', device_serial='', started_at='',
|
||||
last_error='recovered after agent restart', not_before=?
|
||||
WHERE status='running'`, nowStr)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
return int(n), nil
|
||||
}
|
||||
|
||||
// NextDelayed returns the soonest queued job that is still waiting on not_before.
|
||||
func (s *Store) NextDelayed() (Job, time.Duration, bool, error) {
|
||||
nowStr := now()
|
||||
row := s.DB.QueryRow(`
|
||||
SELECT id,stream_id,stream_name,app,channel,reason,status,priority,attempts,max_attempts,
|
||||
device_serial,last_error,not_before,created_at,started_at,finished_at
|
||||
FROM jobs
|
||||
WHERE status='queued' AND not_before!='' AND not_before>?
|
||||
ORDER BY not_before ASC LIMIT 1`, nowStr)
|
||||
j, err := scanJob(row)
|
||||
if err == sql.ErrNoRows {
|
||||
return Job{}, 0, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return Job{}, 0, false, err
|
||||
}
|
||||
until := time.Until(j.NotBefore)
|
||||
if until < 0 {
|
||||
until = 0
|
||||
}
|
||||
return j, until, true, nil
|
||||
}
|
||||
|
||||
// ClearBackoff makes a queued job runnable immediately (e.g. after fixing auth).
|
||||
func (s *Store) ClearBackoff(id int64) error {
|
||||
_, err := s.DB.Exec(`UPDATE jobs SET not_before=? WHERE id=? AND status='queued'`, now(), id)
|
||||
return err
|
||||
}
|
||||
|
||||
// ActiveJobForStream returns the queued/running job for a stream, if any.
|
||||
func (s *Store) ActiveJobForStream(streamID int64) (Job, bool, error) {
|
||||
row := s.DB.QueryRow(`
|
||||
SELECT id,stream_id,stream_name,app,channel,reason,status,priority,attempts,max_attempts,
|
||||
device_serial,last_error,not_before,created_at,started_at,finished_at
|
||||
FROM jobs WHERE stream_id=? AND status IN ('queued','running')
|
||||
ORDER BY id DESC LIMIT 1`, streamID)
|
||||
j, err := scanJob(row)
|
||||
if err == sql.ErrNoRows {
|
||||
return Job{}, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return Job{}, false, err
|
||||
}
|
||||
return j, true, nil
|
||||
}
|
||||
149
apps/agent/internal/worker/worker.go
Normal file
149
apps/agent/internal/worker/worker.go
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
package worker
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"drmdecryption/app"
|
||||
"drmdecryption/phonecap"
|
||||
|
||||
"drmdecryption/apps/agent/internal/client"
|
||||
"drmdecryption/apps/agent/internal/device"
|
||||
"drmdecryption/apps/agent/internal/queue"
|
||||
)
|
||||
|
||||
type Worker struct {
|
||||
ID int
|
||||
API *client.Client
|
||||
Queue *queue.Store
|
||||
Pool *device.Pool
|
||||
AgentID string
|
||||
CaptureWait time.Duration
|
||||
Python string
|
||||
WVD string
|
||||
UserAgent string
|
||||
Log *log.Logger
|
||||
|
||||
lastDelayLog time.Time
|
||||
}
|
||||
|
||||
func (w *Worker) Loop(stop <-chan struct{}) {
|
||||
t := time.NewTicker(2 * time.Second)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
case <-t.C:
|
||||
w.tick()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *Worker) tick() {
|
||||
job, ok, err := w.Queue.ClaimNext()
|
||||
if err != nil {
|
||||
w.logf("claim next: %v", err)
|
||||
return
|
||||
}
|
||||
if !ok {
|
||||
if j, until, found, err := w.Queue.NextDelayed(); err == nil && found {
|
||||
if time.Since(w.lastDelayLog) > 30*time.Second {
|
||||
w.lastDelayLog = time.Now()
|
||||
w.logf("no runnable jobs; next is job=%d stream=%s in %s (backoff)",
|
||||
j.ID, j.StreamName, until.Round(time.Second))
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
w.logf("job=%d running stream=%s app=%s channel=%s reason=%s",
|
||||
job.ID, job.StreamName, job.App, job.Channel, job.Reason)
|
||||
if err := w.runJob(job); err != nil {
|
||||
if device.IsWait(err) || isAuthErr(err) || isUnknownAppErr(err) {
|
||||
// Config/environment issues — do not burn attempts or multi-minute backoff.
|
||||
w.logf("job=%d waiting: %v", job.ID, err)
|
||||
_ = w.Queue.MarkFailed(job.ID, err.Error(), true)
|
||||
return
|
||||
}
|
||||
w.logf("job=%d failed: %v", job.ID, err)
|
||||
_ = w.Queue.MarkFailed(job.ID, err.Error(), false)
|
||||
return
|
||||
}
|
||||
_ = w.Queue.MarkOK(job.ID)
|
||||
w.logf("job=%d ok stream=%s", job.ID, job.StreamName)
|
||||
}
|
||||
|
||||
func isAuthErr(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
s := strings.ToLower(err.Error())
|
||||
return strings.Contains(s, "401") || strings.Contains(s, "unauthorized")
|
||||
}
|
||||
|
||||
func isUnknownAppErr(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
s := strings.ToLower(err.Error())
|
||||
return strings.Contains(s, "unknown app")
|
||||
}
|
||||
|
||||
func (w *Worker) runJob(job queue.Job) error {
|
||||
a, err := app.Get(job.App)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
label := fmt.Sprintf("job-%d:%s", job.ID, job.StreamName)
|
||||
c, serial, err := w.Pool.Acquire(a.Package(), label)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer w.Pool.Release(serial)
|
||||
_ = w.Queue.SetDevice(job.ID, serial)
|
||||
w.logf("job=%d device=%s package=%s", job.ID, serial, a.Package())
|
||||
|
||||
// Short TTL; /ws/agent heartbeats renew while this process is alive.
|
||||
if err := w.API.Claim(job.StreamID, w.AgentID, 90); err != nil {
|
||||
return fmt.Errorf("streamd claim: %w", err)
|
||||
}
|
||||
defer func() { _ = w.API.ReleaseClaim(job.StreamID, w.AgentID) }()
|
||||
|
||||
res, err := phonecap.Run(phonecap.Options{
|
||||
App: a,
|
||||
Channel: job.Channel,
|
||||
Client: c,
|
||||
Python: w.Python,
|
||||
WVD: w.WVD,
|
||||
UserAgent: w.UserAgent,
|
||||
Wait: w.CaptureWait,
|
||||
CloseApp: true,
|
||||
ClearProxy: true,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cred := client.Credentials{
|
||||
MPD: res.MPD,
|
||||
Key: res.Key,
|
||||
PSSH: res.Session.PSSH,
|
||||
Auth: res.Session.Auth,
|
||||
PID: res.Session.PID,
|
||||
}
|
||||
if err := w.API.PostCredentials(job.StreamID, cred); err != nil {
|
||||
return fmt.Errorf("post credentials: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (w *Worker) logf(format string, args ...any) {
|
||||
prefix := fmt.Sprintf("worker/%d: ", w.ID)
|
||||
if w.Log != nil {
|
||||
w.Log.Printf(prefix+format, args...)
|
||||
return
|
||||
}
|
||||
log.Printf(prefix+format, args...)
|
||||
}
|
||||
148
apps/agent/internal/wsclient/wsclient.go
Normal file
148
apps/agent/internal/wsclient/wsclient.go
Normal file
|
|
@ -0,0 +1,148 @@
|
|||
package wsclient
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// Client keeps a streamd /ws/agent connection alive and sends heartbeats
|
||||
// so claims auto-renew; on disconnect streamd releases this agent's claims.
|
||||
type Client struct {
|
||||
BaseURL string
|
||||
Token string
|
||||
AgentID string
|
||||
Log *log.Logger
|
||||
|
||||
stop chan struct{}
|
||||
}
|
||||
|
||||
func New(baseURL, token, agentID string, lg *log.Logger) *Client {
|
||||
if lg == nil {
|
||||
lg = log.Default()
|
||||
}
|
||||
return &Client{
|
||||
BaseURL: strings.TrimRight(baseURL, "/"),
|
||||
Token: token,
|
||||
AgentID: agentID,
|
||||
Log: lg,
|
||||
stop: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Stop() { close(c.stop) }
|
||||
|
||||
// Loop reconnects forever until Stop.
|
||||
func (c *Client) Loop() {
|
||||
backoff := time.Second
|
||||
for {
|
||||
select {
|
||||
case <-c.stop:
|
||||
return
|
||||
default:
|
||||
}
|
||||
if err := c.session(); err != nil {
|
||||
c.Log.Printf("ws: session ended: %v — retry in %s", err, backoff)
|
||||
}
|
||||
select {
|
||||
case <-c.stop:
|
||||
return
|
||||
case <-time.After(backoff):
|
||||
}
|
||||
if backoff < 30*time.Second {
|
||||
backoff *= 2
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) session() error {
|
||||
u, err := url.Parse(c.BaseURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch u.Scheme {
|
||||
case "https":
|
||||
u.Scheme = "wss"
|
||||
default:
|
||||
u.Scheme = "ws"
|
||||
}
|
||||
u.Path = "/ws/agent"
|
||||
q := u.Query()
|
||||
q.Set("agent_id", c.AgentID)
|
||||
if c.Token != "" {
|
||||
q.Set("token", c.Token)
|
||||
}
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
hdr := http.Header{}
|
||||
if c.Token != "" {
|
||||
hdr.Set("Authorization", "Bearer "+c.Token)
|
||||
}
|
||||
dialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second}
|
||||
conn, _, err := dialer.Dial(u.String(), hdr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("dial: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
c.Log.Printf("ws: connected to %s as %s", u.Host, c.AgentID)
|
||||
|
||||
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||
return nil
|
||||
})
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
for {
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(60 * time.Second))
|
||||
var msg map[string]any
|
||||
if json.Unmarshal(data, &msg) == nil {
|
||||
if t, _ := msg["type"].(string); t == "heartbeat_ack" {
|
||||
if ok, _ := msg["ok"].(bool); !ok {
|
||||
c.Log.Printf("ws: heartbeat nack: %v", msg["error"])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
ticker := time.NewTicker(20 * time.Second)
|
||||
defer ticker.Stop()
|
||||
// Immediate heartbeat so existing claims renew right away on reconnect.
|
||||
if err := c.writeHeartbeat(conn); err != nil {
|
||||
return err
|
||||
}
|
||||
for {
|
||||
select {
|
||||
case <-c.stop:
|
||||
_ = conn.WriteJSON(map[string]string{"type": "release_all"})
|
||||
return nil
|
||||
case <-done:
|
||||
return fmt.Errorf("read closed")
|
||||
case <-ticker.C:
|
||||
if err := c.writeHeartbeat(conn); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) writeHeartbeat(conn *websocket.Conn) error {
|
||||
_ = conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
return conn.WriteJSON(map[string]string{
|
||||
"type": "heartbeat",
|
||||
"agent_id": c.AgentID,
|
||||
})
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue