package worker import ( "bytes" "context" "errors" "io" "net" "os" "path/filepath" "sync" "testing" "time" "github.com/kami/maven/internal/audio" ) // stubTranscriber returns a fixed string; satisfies worker.Transcriber. type stubTranscriber struct { mu sync.Mutex gotLast audio.Audio } func (s *stubTranscriber) Transcribe(_ context.Context, req TranscribeReq) (TranscribeResp, error) { s.mu.Lock() defer s.mu.Unlock() s.gotLast = req.Audio return TranscribeResp{Text: "hello from stt", Confidence: 0.9}, nil } type errTranscriber struct{} func (errTranscriber) Transcribe(context.Context, TranscribeReq) (TranscribeResp, error) { return TranscribeResp{}, errors.New("synth failed") } type stubSynthesizer struct{} func (stubSynthesizer) Synthesize(_ context.Context, req SynthesizeReq) (SynthesizeResp, error) { pcm := make([]byte, 3200) // 100ms of silence @ 16k mono int16 _ = req return SynthesizeResp{Audio: audio.Audio{Format: audio.PCM16kMono, Bytes: pcm}}, nil } // newServer builds a Server on a temp socket, starts Serve in a goroutine, // returns the Server + path + a cleanup. Tests use this to get a real // round-trip over a unix socket. func newServer(t *testing.T, srv *Server) (*Server, string, func()) { t.Helper() dir := t.TempDir() sock := filepath.Join(dir, "worker.sock") srv.path = sock if err := srv.Listen(); err != nil { t.Fatalf("listen: %v", err) } go func() { if err := srv.Serve(); err != nil { t.Logf("serve ended: %v", err) } }() cleanup := func() { _ = srv.Close() } return srv, sock, cleanup } func TestServerRejectsUnknownMethod(t *testing.T) { t.Parallel() srv := NewSynthesizerServer("", stubSynthesizer{}) _, sock, cleanup := newServer(t, srv) defer cleanup() c := Dial(sock) defer c.Close() _, err := c.Transcribe(context.Background(), TranscribeReq{Audio: audio.Audio{Format: audio.PCM16kMono, Bytes: []byte{1, 2}}}) if err == nil { t.Fatalf("Transcribe on a tts server should error") } if !errors.Is(err, ErrUnknownMethod) { t.Fatalf("err should be ErrUnknownMethod, got: %v", err) } } func TestServerTranscribeRoundTrip(t *testing.T) { t.Parallel() tr := &stubTranscriber{} srv := NewServer("", tr) _, sock, cleanup := newServer(t, srv) defer cleanup() c := Dial(sock) defer c.Close() in := audio.Audio{Format: audio.PCM16kMono, Bytes: []byte("hello audio payload")} resp, err := c.Transcribe(context.Background(), TranscribeReq{Audio: in, Lang: "ru"}) if err != nil { t.Fatalf("Transcribe: %v", err) } if resp.Text != "hello from stt" { t.Fatalf("Text: %q", resp.Text) } if resp.Confidence != 0.9 { t.Fatalf("Confidence: %v", resp.Confidence) } if !bytes.Equal(tr.gotLast.Bytes, in.Bytes) { t.Fatalf("audio bytes did not round-trip: in=%v got=%v", in.Bytes, tr.gotLast.Bytes) } } func TestServerSynthesizeRoundTrip(t *testing.T) { t.Parallel() srv := NewSynthesizerServer("", stubSynthesizer{}) _, sock, cleanup := newServer(t, srv) defer cleanup() c := Dial(sock) defer c.Close() resp, err := c.Synthesize(context.Background(), SynthesizeReq{Text: "привет", Lang: "ru"}) if err != nil { t.Fatalf("Synthesize: %v", err) } if !resp.Audio.Format.IsValid() { t.Fatalf("audio format invalid: %+v", resp.Audio.Format) } if len(resp.Audio.Bytes) != 3200 { t.Fatalf("audio bytes len: %d, want 3200", len(resp.Audio.Bytes)) } } func TestServerUnsetVerbReturnsUnknownMethod(t *testing.T) { t.Parallel() srv := NewServer("", &stubTranscriber{}) // transcriber-only _, sock, cleanup := newServer(t, srv) defer cleanup() c := Dial(sock) defer c.Close() _, err := c.Synthesize(context.Background(), SynthesizeReq{Text: "x"}) if !errors.Is(err, ErrUnknownMethod) { t.Fatalf("Synthesize on a stt-only server: want ErrUnknownMethod, got %v", err) } } func TestServerDispatchesBadParams(t *testing.T) { t.Parallel() srv := NewServer("", &stubTranscriber{}) _, sock, cleanup := newServer(t, srv) defer cleanup() // raw socket: ship a Request frame with malformed params JSON. We can't // use writeFrame because Request marshals RawMessage and rejects invalid // JSON itself — so build the bytes by hand: a valid Request envelope // whose `p` field is a syntactically broken JSON string. conn, err := net.Dial("unix", sock) if err != nil { t.Fatalf("dial: %v", err) } defer conn.Close() // {"m":"transcribe","p":{not json}} — but the `p` value is invalid JSON. // We instead send `{"m":"transcribe","p":""}` — a valid // JSON frame whose params unmarshal fails into TranscribeReq (string // into a struct). That path hits unmarshalParams' error ⇒ codeBadParams. body := []byte(`{"m":"transcribe","p":"not an object"}`) var hdr [4]byte hdr[0] = byte(len(body) >> 24) hdr[1] = byte(len(body) >> 16) hdr[2] = byte(len(body) >> 8) hdr[3] = byte(len(body)) if _, err := conn.Write(hdr[:]); err != nil { t.Fatalf("write hdr: %v", err) } if _, err := conn.Write(body); err != nil { t.Fatalf("write body: %v", err) } var resp Response if err := readFrame(conn, &resp); err != nil { t.Fatalf("readFrame: %v", err) } if resp.Error == nil || resp.Error.Code != codeBadParams { t.Fatalf("want codeBadParams, got %+v", resp.Error) } } func TestServerInternalErrorPropagates(t *testing.T) { t.Parallel() srv := NewServer("", errTranscriber{}) _, sock, cleanup := newServer(t, srv) defer cleanup() c := Dial(sock) defer c.Close() _, err := c.Transcribe(context.Background(), TranscribeReq{Audio: audio.Audio{Format: audio.PCM16kMono}}) if err == nil { t.Fatalf("want error from errTranscriber") } if errors.Is(err, ErrUnknownMethod) || errors.Is(err, ErrBadParams) { t.Fatalf("internal error should not match known sentinels: %v", err) } if !contains(err.Error(), "synth failed") { t.Fatalf("err should carry internal message text, got: %v", err) } } func TestClientContextCancelStopsCall(t *testing.T) { t.Parallel() // a server that never replies — transcriber blocked on a chan. hang := &hangTranscriber{ok: make(chan struct{})} srv := NewServer("", hang) _, sock, cleanup := newServer(t, srv) defer cleanup() defer close(hang.ok) c := Dial(sock) defer c.Close() ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() _, err := c.Transcribe(ctx, TranscribeReq{Audio: audio.Audio{Format: audio.PCM16kMono, Bytes: []byte("x")}}) if err == nil { t.Fatalf("want ctx timeout error") } if !contains(err.Error(), "dial") && !errors.Is(err, context.DeadlineExceeded) { // dial may succeed within 100ms; if so, the deadline set on the conn // surfaces as an i/o timeout from writeFrame/readFrame. Either path // is acceptable; we just assert the call returned in finite time. } } type hangTranscriber struct { ok chan struct{} } func (h *hangTranscriber) Transcribe(ctx context.Context, _ TranscribeReq) (TranscribeResp, error) { select { case <-h.ok: return TranscribeResp{}, nil case <-ctx.Done(): return TranscribeResp{}, ctx.Err() } } func TestServerCloseIsIdempotent(t *testing.T) { t.Parallel() srv := NewServer("", &stubTranscriber{}) _, _, cleanup := newServer(t, srv) cleanup() // first close if err := srv.Close(); err != nil { t.Fatalf("second Close should be no-op, got: %v", err) } } func TestFrameTooLargeRejected(t *testing.T) { t.Parallel() // build a frame whose length prefix exceeds maxFrame; ensure readFrame // returns ErrFrameTooLarge. r, w := io.Pipe() go func() { var hdr [4]byte hdr[0] = 0xff hdr[1] = 0xff hdr[2] = 0xff hdr[3] = 0xff _, _ = w.Write(hdr[:]) _ = w.Close() }() var v any err := readFrame(r, &v) if !errors.Is(err, ErrFrameTooLarge) { t.Fatalf("want ErrFrameTooLarge, got %v", err) } _ = r.Close() } func TestPathAfterListen(t *testing.T) { t.Parallel() srv := NewServer("", &stubTranscriber{}) dir := t.TempDir() sock := filepath.Join(dir, "x.sock") srv.path = sock if err := srv.Listen(); err != nil { t.Fatalf("listen: %v", err) } defer srv.Close() if srv.Path() != sock { t.Fatalf("Path: got %q, want %q", srv.Path(), sock) } if _, err := os.Stat(sock); err != nil { t.Fatalf("socket file missing after listen: %v", err) } } func contains(haystack, needle string) bool { return len(haystack) >= len(needle) && (bytes.Contains([]byte(haystack), []byte(needle))) }