From cdbe2c250fb275b50ebf75b7149d801160b521cf Mon Sep 17 00:00:00 2001 From: Kieran Klukas Date: Sun, 24 May 2026 17:57:23 -0400 Subject: [PATCH] feat: use streaming to auto reload --- server/internal/api/web/chat.go | 98 +++++++++++++++---- server/internal/api/web/web.go | 67 +++++++++++-- web/src/lib/stream.ts | 9 +- web/src/routes/chat/+page.svelte | 155 ++++++++++++++++++++++++++++++- 4 files changed, 294 insertions(+), 35 deletions(-) diff --git a/server/internal/api/web/chat.go b/server/internal/api/web/chat.go index 4e421cc..6d4eecd 100644 --- a/server/internal/api/web/chat.go +++ b/server/internal/api/web/chat.go @@ -186,6 +186,17 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { CreatedAt: now + 1, }) + streamID := uuid.NewString() + _, _ = s.Q.CreateStream(r.Context(), store.CreateStreamParams{ + ID: streamID, + ConversationID: convID, + UserID: u.ID, + AssistantMessageID: sql.NullString{String: assistantMsgID, Valid: true}, + IdempotencyKey: uuid.NewString(), + Model: req.Model, + StartedAt: now, + }) + // Pick provider and upstream model name. upstreamModel := req.Model var pc *provider.Client @@ -237,27 +248,54 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { w.WriteHeader(200) flusher, _ := w.(http.Flusher) - emit := func(v any) { - b, _ := json.Marshal(v) + var full strings.Builder + var usage *provider.Usage + + _ = s.Q.SetStreamStatus(r.Context(), store.SetStreamStatusParams{ + Status: "running", + FinishedAt: sql.NullInt64{}, + ErrorCode: sql.NullString{}, + ErrorMessage: sql.NullString{}, + ID: streamID, + }) + + seq := int64(0) + clientGone := false + ctxDone := r.Context().Done() + genCtx := context.Background() + + emit := func(event string, payload map[string]any) { + seq++ + payload["seq"] = seq + b, _ := json.Marshal(payload) + _ = s.Q.AppendStreamChunk(genCtx, store.AppendStreamChunkParams{ + StreamID: streamID, + Seq: seq, + Event: event, + Data: string(b), + CreatedAt: time.Now().Unix(), + }) + if clientGone { + return + } _, _ = fmt.Fprintf(w, "data: %s\n\n", b) if flusher != nil { flusher.Flush() } } - emit(map[string]any{ + startPayload := map[string]any{ "type": "start", "conversation_id": convID, "user_message_id": userMsgID, "assistant_message_id": assistantMsgID, - }) - - var full strings.Builder - var usage *provider.Usage + "stream_id": streamID, + } + emit("start", startPayload) const maxIter = 5 for iter := 0; iter < maxIter; iter++ { - chunks, errs, err := pc.StreamChat(r.Context(), provider.ChatRequest{ + chunks, errs, err := pc.StreamChat(genCtx, provider.ChatRequest{ Model: upstreamModel, Messages: provMsgs, Tools: tools.Definitions(), @@ -276,11 +314,10 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { streamDone := false for !streamDone { select { - case <-r.Context().Done(): - if !isFree && reqID != "" { - go finishWebReq(s.Q, reqID, 0, 0, 0, "canceled") - } - return + case <-ctxDone: + clientGone = true + ctxDone = nil + continue case ch, ok := <-chunks: if ch.Usage != nil { usage = ch.Usage @@ -291,7 +328,7 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { } if ch.Delta != "" { full.WriteString(ch.Delta) - emit(map[string]any{"type": "delta", "content": ch.Delta}) + emit("delta", map[string]any{"type": "delta", "content": ch.Delta}) } for _, tcd := range ch.ToolCalls { if tcd.Index < 0 { @@ -319,7 +356,14 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { } case e := <-errs: if e != nil { - emit(map[string]any{"type": "error", "message": e.Error()}) + emit("error", map[string]any{"type": "error", "message": e.Error()}) + _ = s.Q.SetStreamStatus(genCtx, store.SetStreamStatusParams{ + Status: "error", + FinishedAt: sql.NullInt64{Int64: time.Now().Unix(), Valid: true}, + ErrorCode: sql.NullString{String: "provider_down", Valid: true}, + ErrorMessage: sql.NullString{String: e.Error(), Valid: true}, + ID: streamID, + }) if !isFree && reqID != "" { go finishWebReq(s.Q, reqID, 0, 0, 0, "error") } @@ -331,7 +375,7 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { // Tool invocation: execute tools, emit events, extend messages, re-stream. if finishReason == "tool_calls" && len(toolCalls) > 0 { for _, tc := range toolCalls { - emit(map[string]any{ + emit("tool_call", map[string]any{ "type": "tool_call", "id": tc.ID, "name": tc.Function.Name, @@ -352,11 +396,11 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { // Execute each tool and append a tool-role message with the result. for _, tc := range toolCalls { - result, toolErr := tools.Execute(r.Context(), s.Q, u.ID, tc.Function.Name, tc.Function.Arguments) + result, toolErr := tools.Execute(genCtx, s.Q, u.ID, tc.Function.Name, tc.Function.Arguments) if toolErr != nil { result = fmt.Sprintf("error: %v", toolErr) } - emit(map[string]any{ + emit("tool_result", map[string]any{ "type": "tool_result", "tool_call_id": tc.ID, "content": result, @@ -379,7 +423,14 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { "total_tokens": usage.TotalTokens, } } - emit(payload) + emit("done", payload) + _ = s.Q.SetStreamStatus(genCtx, store.SetStreamStatusParams{ + Status: "done", + FinishedAt: sql.NullInt64{Int64: time.Now().Unix(), Valid: true}, + ErrorCode: sql.NullString{}, + ErrorMessage: sql.NullString{}, + ID: streamID, + }) go s.finalizeChatMsg(convID, assistantMsgID, full.String(), isFree, sel, reqID, usage) return } @@ -393,7 +444,14 @@ func (s *Server) handleChat(w http.ResponseWriter, r *http.Request) { "total_tokens": usage.TotalTokens, } } - emit(payload) + emit("done", payload) + _ = s.Q.SetStreamStatus(genCtx, store.SetStreamStatusParams{ + Status: "done", + FinishedAt: sql.NullInt64{Int64: time.Now().Unix(), Valid: true}, + ErrorCode: sql.NullString{}, + ErrorMessage: sql.NullString{}, + ID: streamID, + }) go s.finalizeChatMsg(convID, assistantMsgID, full.String(), isFree, sel, reqID, usage) } diff --git a/server/internal/api/web/web.go b/server/internal/api/web/web.go index 8a2e6ce..915e53a 100644 --- a/server/internal/api/web/web.go +++ b/server/internal/api/web/web.go @@ -10,8 +10,8 @@ package web import ( "context" "encoding/json" - "errors" "net/http" + "strconv" "strings" "time" @@ -485,21 +485,24 @@ func (s *Server) handleRevokeKey(w http.ResponseWriter, r *http.Request) { // ---- chat / streams ---------------------------------------------------- // handleStreamEvents serves SSE events for a given stream id, optionally -// resuming from ?after_seq=N. Stub: replay from DB only. +// resuming from ?after_seq=N. It replays durable chunks from DB first, then +// attaches to the in-memory bus for live events. func (s *Server) handleStreamEvents(w http.ResponseWriter, r *http.Request) { streamID := chi.URLParam(r, "id") + afterSeq := int64(0) + if v := r.URL.Query().Get("after_seq"); v != "" { + if n, err := strconv.ParseInt(v, 10, 64); err == nil && n >= 0 { + afterSeq = n + } + } + w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.Header().Set("Connection", "keep-alive") w.WriteHeader(200) flusher, _ := w.(http.Flusher) - events, err := stream.Replay(r.Context(), s.Q, streamID, 0) - if err != nil && !errors.Is(err, http.ErrAbortHandler) { - // fall through; partial replay is fine - _ = err - } - for _, ev := range events { + emit := func(ev stream.Event) { b, _ := json.Marshal(ev) _, _ = w.Write([]byte("data: ")) _, _ = w.Write(b) @@ -508,4 +511,52 @@ func (s *Server) handleStreamEvents(w http.ResponseWriter, r *http.Request) { flusher.Flush() } } + + // 1) Replay durable chunks after the requested seq. + events, err := stream.Replay(r.Context(), s.Q, streamID, afterSeq) + if err == nil { + for _, ev := range events { + emit(ev) + afterSeq = ev.Seq + if ev.Type == "done" || ev.Type == "error" { + return + } + } + } + + // 2) Attach to live bus for tailing events. + bus := s.Hub.Subscriber(streamID) + ch, done, doneEv := bus.Subscribe(64) + if done { + if doneEv.Seq > afterSeq { + emit(doneEv) + } + return + } + + heartbeat := time.NewTicker(15 * time.Second) + defer heartbeat.Stop() + for { + select { + case <-r.Context().Done(): + return + case <-heartbeat.C: + _, _ = w.Write([]byte(": ping\n\n")) + if flusher != nil { + flusher.Flush() + } + case ev, ok := <-ch: + if !ok { + return + } + if ev.Seq <= afterSeq { + continue + } + emit(ev) + afterSeq = ev.Seq + if ev.Type == "done" || ev.Type == "error" { + return + } + } + } } diff --git a/web/src/lib/stream.ts b/web/src/lib/stream.ts index 4811a1d..d325ab2 100644 --- a/web/src/lib/stream.ts +++ b/web/src/lib/stream.ts @@ -11,10 +11,11 @@ export interface StreamEvent { seq: number; - type: 'delta' | 'usage' | 'error' | 'done'; + type: 'delta' | 'usage' | 'error' | 'done' | 'tool_call' | 'tool_result' | 'start'; content?: string; usage?: { input_tokens: number; output_tokens: number }; error?: { code: string; message: string }; + [extra: string]: any; } export interface StreamHandlers { @@ -22,15 +23,15 @@ export interface StreamHandlers { onClose?(reason: 'done' | 'error' | 'aborted'): void; } -export function consume(streamID: string, handlers: StreamHandlers): () => void { +export function consume(streamID: string, handlers: StreamHandlers, initialAfterSeq = 0): () => void { let aborted = false; - let lastSeq = 0; + let lastSeq = initialAfterSeq; let controller = new AbortController(); const loop = async () => { while (!aborted) { try { - const res = await fetch(`/api/streams/${streamID}/events?after_seq=${lastSeq}`, { + const res = await fetch(`/api/streams/${streamID}/events?after_seq=${Math.max(0, lastSeq)}`, { credentials: 'include', signal: controller.signal }); diff --git a/web/src/routes/chat/+page.svelte b/web/src/routes/chat/+page.svelte index a44b3ef..011bf09 100644 --- a/web/src/routes/chat/+page.svelte +++ b/web/src/routes/chat/+page.svelte @@ -1,5 +1,5 @@