probe.go
4463 bytes
1package services
2
3import (
4 "bytes"
5 "context"
6 "encoding/binary"
7 "fmt"
8 "io"
9 "math"
10 "net/http"
11 "strings"
12 "sync"
13 "time"
14
15 "git.theedgeofrage.com/TheEdgeOfRage/kaiwari/internal/stt"
16 "git.theedgeofrage.com/TheEdgeOfRage/kaiwari/internal/tts"
17)
18
19const (
20 probeTimeout = 10 * time.Second
21 pollInterval = 2 * time.Second
22)
23
24// waitHealthy polls the TTS, STT, and LLM probes concurrently until all three
25// succeed or the total readiness deadline (or ctx) expires. Each probe is a real
26// request round-trip through the same client methods the app uses.
27func (m *Manager) waitHealthy(ctx context.Context) error {
28 ctx, cancel := context.WithDeadline(ctx, time.Now().Add(modelReadyTimeout))
29 defer cancel()
30
31 hc := newProbeHTTP()
32 ttsClient := tts.NewClient(httpURL(AudioListen), hc)
33 sttClient := stt.NewASRClient(httpURL(AudioListen), hc)
34 llmURL := httpURL(LLMListen) + "/v1/chat/completions"
35 audio := probeWAV()
36
37 var failed []string
38 for {
39 failed = probeAll(ctx, ttsClient, sttClient, llmURL, hc, audio)
40 if len(failed) == 0 {
41 return nil
42 }
43 if !sleepCtx(ctx, pollInterval) {
44 break
45 }
46 }
47 return fmt.Errorf("services: model services not healthy after %s (unhealthy: %s)", modelReadyTimeout, strings.Join(failed, ", "))
48}
49
50func probeAll(ctx context.Context, ttsClient *tts.Client, sttClient *stt.ASRClient, llmURL string, hc *http.Client, audio []byte) []string {
51 names := []string{"tts", "stt", "llm"}
52 probes := []func(context.Context) error{
53 func(ctx context.Context) error { _, e := ttsClient.Speech(ctx, "あ", tts.DefaultVoice); return e },
54 func(ctx context.Context) error { _, e := sttClient.TranscribeBytes(ctx, audio); return e },
55 func(ctx context.Context) error { return probeLLMCompletion(ctx, llmURL, hc) },
56 }
57
58 errs := make([]error, len(names))
59 var wg sync.WaitGroup
60 for i := range probes {
61 wg.Add(1)
62 go func(i int) {
63 defer wg.Done()
64 errs[i] = probes[i](ctx)
65 }(i)
66 }
67 wg.Wait()
68
69 var failed []string
70 for i, e := range errs {
71 if e != nil {
72 failed = append(failed, names[i])
73 }
74 }
75 return failed
76}
77
78func probeLLMCompletion(ctx context.Context, url string, hc *http.Client) error {
79 body := `{"messages":[{"role":"user","content":"Ready."}],"max_tokens":1,"stream":false}`
80 req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(body))
81 if err != nil {
82 return fmt.Errorf("services: llm probe build request: %w", err)
83 }
84 req.Header.Set("Content-Type", "application/json")
85 resp, err := hc.Do(req)
86 if err != nil {
87 return fmt.Errorf("services: llm probe request to %s: %w", url, err)
88 }
89 defer func() { _ = resp.Body.Close() }()
90 detail, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
91 if resp.StatusCode < 200 || resp.StatusCode >= 300 {
92 return fmt.Errorf("services: llm probe %s returned HTTP %d: %s", url, resp.StatusCode, strings.TrimSpace(string(detail)))
93 }
94 return nil
95}
96
97func newProbeHTTP() *http.Client { return &http.Client{Timeout: probeTimeout} }
98
99func httpURL(hostPort string) string { return "http://" + hostPort }
100
101func sleepCtx(ctx context.Context, d time.Duration) bool {
102 t := time.NewTimer(d)
103 defer t.Stop()
104 select {
105 case <-ctx.Done():
106 return false
107 case <-t.C:
108 return true
109 }
110}
111
112// probeWAV builds a short low-amplitude sine WAV (16 kHz mono S16_LE) in memory
113// for the ASR upload so the model has real audio to transcribe.
114func probeWAV() []byte {
115 const sampleRate = 16000
116 samples := int(float64(sampleRate) * 0.5)
117 data := make([]byte, samples*2)
118 for i := 0; i < samples; i++ {
119 v := int16(800 * math.Sin(2*math.Pi*440*float64(i)/float64(sampleRate)))
120 binary.LittleEndian.PutUint16(data[i*2:], uint16(v))
121 }
122 var buf bytes.Buffer
123 buf.Write(wavHeader(len(data), sampleRate))
124 buf.Write(data)
125 return buf.Bytes()
126}
127
128func wavHeader(dataSize, sampleRate int) []byte {
129 var buf bytes.Buffer
130 writeU32 := func(v uint32) { var b [4]byte; binary.LittleEndian.PutUint32(b[:], v); buf.Write(b[:]) }
131 writeU16 := func(v uint16) { var b [2]byte; binary.LittleEndian.PutUint16(b[:], v); buf.Write(b[:]) }
132
133 buf.WriteString("RIFF")
134 writeU32(36 + uint32(dataSize))
135 buf.WriteString("WAVE")
136 buf.WriteString("fmt ")
137 writeU32(16) // PCM fmt chunk size
138 writeU16(1) // PCM
139 writeU16(1) // mono
140 writeU32(uint32(sampleRate))
141 writeU32(uint32(sampleRate) * 2) // byte rate
142 writeU16(2) // block align
143 writeU16(16) // bits per sample
144 buf.WriteString("data")
145 writeU32(uint32(dataSize))
146 return buf.Bytes()
147}