Parent directory

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}