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}