Parent directory

handlers.go

7613 bytes
  1package server
  2
  3import (
  4	"context"
  5	"encoding/base64"
  6	"encoding/binary"
  7	"encoding/json"
  8	"errors"
  9	"io"
 10	"log/slog"
 11	"net/http"
 12	"strings"
 13
 14	"git.theedgeofrage.com/TheEdgeOfRage/kaiwari/internal/game"
 15)
 16
 17type judgeWire struct {
 18	Score    int    `json:"score"`
 19	Romaji   string `json:"romaji,omitempty"`
 20	Feedback string `json:"feedback,omitempty"`
 21}
 22
 23type askWire struct {
 24	Romaji      string `json:"romaji,omitempty"`
 25	Translation string `json:"translation,omitempty"`
 26	Breakdown   string `json:"breakdown,omitempty"`
 27}
 28
 29// logEntry is one display entry on the wire. Kana never crosses the protocol.
 30type logEntry struct {
 31	Action    string     `json:"action,omitempty"`
 32	Desc      string     `json:"desc,omitempty"`
 33	Romaji    string     `json:"romaji,omitempty"`
 34	English   string     `json:"english,omitempty"`
 35	HasSpeech bool       `json:"hasSpeech"`
 36	Judge     *judgeWire `json:"judge,omitempty"`
 37	Ask       *askWire   `json:"ask,omitempty"`
 38}
 39
 40func toLogEntry(e game.DisplayEntry) logEntry {
 41	out := logEntry{
 42		Action:    e.Action,
 43		Desc:      e.Desc,
 44		Romaji:    e.Romaji,
 45		English:   e.English,
 46		HasSpeech: e.HasSpeech,
 47	}
 48	if e.HasJudge {
 49		out.Judge = &judgeWire{Score: e.Score, Romaji: e.PlayerRomaji, Feedback: e.Feedback}
 50	}
 51	if e.HasAsk {
 52		out.Ask = &askWire{Romaji: e.AskRomaji, Translation: e.AskTranslation, Breakdown: e.AskBreakdown}
 53	}
 54	return out
 55}
 56
 57type turnResponse struct {
 58	logEntry
 59	Location   string `json:"location,omitempty"`
 60	Talk       string `json:"talk,omitempty"`
 61	Audio      string `json:"audio,omitempty"`
 62	SpeakError string `json:"speakError,omitempty"`
 63	JudgeError string `json:"judgeError,omitempty"`
 64}
 65
 66type flashcardResponse struct {
 67	Added int `json:"added"`
 68}
 69
 70type errorWire struct {
 71	Error string `json:"error"`
 72}
 73
 74func writeJSON(w http.ResponseWriter, status int, v any) {
 75	w.Header().Set("Content-Type", "application/json")
 76	w.WriteHeader(status)
 77	_ = json.NewEncoder(w).Encode(v)
 78}
 79
 80func writeError(w http.ResponseWriter, status int, msg string) {
 81	writeJSON(w, status, errorWire{Error: msg})
 82}
 83
 84// createSessionResponse is the start-adventure response: the new session id plus
 85// the opening turn (the model's description of the player's starting area).
 86type createSessionResponse struct {
 87	ID string `json:"id"`
 88	turnResponse
 89}
 90
 91func (s *Server) createSession(w http.ResponseWriter, r *http.Request) {
 92	sess := s.newSession()
 93
 94	sess.mu.Lock()
 95	ctx, cancel := context.WithTimeout(r.Context(), turnTimeout)
 96	res := sess.orch.StartAdventure(ctx)
 97	cancel()
 98	if res.Err != nil {
 99		sess.mu.Unlock()
100		slog.Error("start adventure failed", "session", sess.id, "err", res.Err)
101		writeError(w, http.StatusBadGateway, res.Err.Error())
102		return
103	}
104	out := createSessionResponse{ID: sess.id, turnResponse: buildTurnResponse(sess, res)}
105	sess.mu.Unlock()
106
107	s.mu.Lock()
108	s.sessions[sess.id] = sess
109	s.mu.Unlock()
110
111	writeJSON(w, http.StatusOK, out)
112}
113
114func (s *Server) getLog(w http.ResponseWriter, r *http.Request) {
115	sess := s.session(r.PathValue("id"))
116	if sess == nil {
117		writeError(w, http.StatusNotFound, "unknown session")
118		return
119	}
120
121	sess.mu.Lock()
122	entries := sess.state.Display()
123	sess.mu.Unlock()
124
125	out := make([]logEntry, 0, len(entries))
126	for _, e := range entries {
127		out = append(out, toLogEntry(e))
128	}
129	writeJSON(w, http.StatusOK, out)
130}
131
132func (s *Server) postAction(w http.ResponseWriter, r *http.Request) {
133	sess := s.session(r.PathValue("id"))
134	if sess == nil {
135		writeError(w, http.StatusNotFound, "unknown session")
136		return
137	}
138
139	var body struct {
140		Action string `json:"action"`
141	}
142	if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
143		writeError(w, http.StatusBadRequest, "invalid action body")
144		return
145	}
146	action := body.Action
147	if strings.TrimSpace(action) == "" {
148		writeError(w, http.StatusBadRequest, "empty action")
149		return
150	}
151
152	sess.mu.Lock()
153	defer sess.mu.Unlock()
154	ctx, cancel := context.WithTimeout(r.Context(), turnTimeout)
155	defer cancel()
156
157	switch {
158	case strings.HasPrefix(action, "?"):
159		if strings.TrimSpace(strings.TrimPrefix(action, "?")) == "" {
160			writeError(w, http.StatusBadRequest, "empty question")
161			return
162		}
163		finishTurn(w, sess, sess.orch.AskTurn(ctx, action))
164	case strings.HasPrefix(action, "!"):
165		if strings.TrimSpace(strings.TrimPrefix(action, "!")) == "" {
166			writeError(w, http.StatusBadRequest, "empty flashcard instructions")
167			return
168		}
169		count, err := sess.orch.GenerateFlashcards(ctx, action)
170		if err != nil {
171			slog.Error("flashcards failed", "session", sess.id, "err", err)
172			writeError(w, http.StatusBadGateway, err.Error())
173			return
174		}
175		writeJSON(w, http.StatusOK, flashcardResponse{Added: count})
176	default:
177		finishTurn(w, sess, sess.orch.ActionTurn(ctx, action))
178	}
179}
180
181func (s *Server) postSpeech(w http.ResponseWriter, r *http.Request) {
182	sess := s.session(r.PathValue("id"))
183	if sess == nil {
184		writeError(w, http.StatusNotFound, "unknown session")
185		return
186	}
187
188	r.Body = http.MaxBytesReader(w, r.Body, maxUpload)
189	wav, err := io.ReadAll(r.Body)
190	if err != nil {
191		var tooLarge *http.MaxBytesError
192		if errors.As(err, &tooLarge) {
193			writeError(w, http.StatusRequestEntityTooLarge, "audio too large")
194			return
195		}
196		writeError(w, http.StatusBadRequest, "invalid audio body")
197		return
198	}
199	if len(wav) < 12 || string(wav[0:4]) != "RIFF" || string(wav[8:12]) != "WAVE" {
200		writeError(w, http.StatusBadRequest, "not a WAV file")
201		return
202	}
203	normalizeWav(wav)
204
205	sess.mu.Lock()
206	defer sess.mu.Unlock()
207	ctx, cancel := context.WithTimeout(r.Context(), turnTimeout)
208	defer cancel()
209
210	sess.input.SetWAV(wav)
211	finishTurn(w, sess, sess.orch.FinishSpeak(ctx))
212}
213
214// normalizeWav repairs header sizes that AVAudioRecorder leaves stale: it
215// preallocates a padded header, streams PCM after it, and never updates the
216// RIFF size or the data chunk size (written as 0). Strict readers reject such
217// files.
218func normalizeWav(wav []byte) {
219	off := 12
220	for off+8 <= len(wav) {
221		id := string(wav[off : off+4])
222		size := int(binary.LittleEndian.Uint32(wav[off+4 : off+8]))
223		if id == "data" {
224			if size != len(wav)-off-8 {
225				binary.LittleEndian.PutUint32(wav[off+4:], uint32(len(wav)-off-8))
226			}
227			break
228		}
229		off += 8 + size + size%2
230	}
231	if int(binary.LittleEndian.Uint32(wav[4:8])) != len(wav)-8 {
232		binary.LittleEndian.PutUint32(wav[4:], uint32(len(wav)-8))
233	}
234}
235
236// buildTurnResponse assembles the wire response for a completed turn from the last
237// recorded display entry. The session lock must be held.
238func buildTurnResponse(sess *session, res game.TurnResult) turnResponse {
239	entries := sess.state.Display()
240	out := turnResponse{logEntry: toLogEntry(entries[len(entries)-1])}
241	out.Location = sess.state.Location()
242	out.Talk = sess.state.Talk()
243	if wav := sess.output.TakeLast(); len(wav) > 0 {
244		out.Audio = base64.StdEncoding.EncodeToString(wav)
245	}
246	if res.SpeakErr != nil {
247		out.SpeakError = res.SpeakErr.Error()
248		slog.Warn("speak failed", "session", sess.id, "err", res.SpeakErr)
249	}
250	if res.JudgeErr != nil {
251		out.JudgeError = res.JudgeErr.Error()
252		slog.Warn("judge failed", "session", sess.id, "err", res.JudgeErr)
253	}
254	return out
255}
256
257// finishTurn writes a successful turn response from the display entry the
258// orchestrator just recorded. The session lock must be held.
259func finishTurn(w http.ResponseWriter, sess *session, res game.TurnResult) {
260	if res.Err != nil {
261		status := http.StatusBadGateway
262		if errors.Is(res.Err, game.ErrEmptyTranscript) {
263			status = http.StatusBadRequest
264		}
265		slog.Error("turn failed", "session", sess.id, "err", res.Err)
266		writeError(w, status, res.Err.Error())
267		return
268	}
269	writeJSON(w, http.StatusOK, buildTurnResponse(sess, res))
270}