Skip to content

Commit 44620e8

Browse files
committed
feat: wire memory into Web UI (serve.go)
- Agent created once per WebSocket connection (buffer continuity) - Buffer restored on session resume - Buffer appended per turn (user + agent) - Buffer persisted to session on save - OnSessionEnd on WebSocket disconnect - Removed placeholder helper functions
1 parent e84ac26 commit 44620e8

2 files changed

Lines changed: 111 additions & 40 deletions

File tree

cmd/kode/serve.go

Lines changed: 102 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,6 @@ import (
1717
"github.com/BackendStack21/kode"
1818
"github.com/BackendStack21/kode/internal/config"
1919
"github.com/BackendStack21/kode/internal/llm"
20-
"github.com/BackendStack21/kode/internal/render"
2120
"github.com/BackendStack21/kode/internal/resource"
2221
"github.com/BackendStack21/kode/internal/session"
2322
"github.com/BackendStack21/kode/internal/skills"
@@ -151,16 +150,17 @@ func newServeAgent(resolved config.ResolvedConfig, system string) (*kode.Agent,
151150
return agent, sandboxCleanup, nil
152151
}
153152

154-
// ── WebSocket Handler ───────────────────────────────────────────────────
153+
// ── WebSocket Types ────────────────────────────────────────────────────
155154

156155
type wsClientMsg struct {
157156
Type string `json:"type"`
158157
Content string `json:"content"`
159158
SessionID string `json:"session_id"`
160159
}
161160

161+
// ── WebSocket Handler ──────────────────────────────────────────────────
162+
162163
func handleWebSocket(store *session.Store, resources *resource.Registry, resolved config.ResolvedConfig, system string) http.HandlerFunc {
163-
// Resolve config once at startup — reconnects use the same config
164164
return func(w http.ResponseWriter, r *http.Request) {
165165
conn, err := ws.Upgrade(w, r)
166166
if err != nil {
@@ -169,13 +169,28 @@ func handleWebSocket(store *session.Store, resources *resource.Registry, resolve
169169
}
170170
defer conn.Close()
171171

172+
// Create ONE agent per WebSocket connection — provides buffer
173+
// continuity across turns within the same session.
174+
agent, sandboxCleanup, err := newServeAgent(resolved, system)
175+
if err != nil {
176+
writeWSError(conn, fmt.Sprintf("agent: %v", err))
177+
return
178+
}
179+
defer agent.Close()
180+
if sandboxCleanup != nil {
181+
defer sandboxCleanup()
182+
}
183+
172184
ctx, cancel := signal.NotifyContext(r.Context(), os.Interrupt)
173185
defer cancel()
174186

187+
// Track the current session across WebSocket messages
188+
var currentSession *session.Session
189+
175190
for {
176191
opcode, data, err := conn.ReadMessage()
177192
if err != nil {
178-
return
193+
break
179194
}
180195
if opcode != ws.OpText {
181196
continue
@@ -191,13 +206,48 @@ func handleWebSocket(store *session.Store, resources *resource.Registry, resolve
191206
continue
192207
}
193208

194-
// Run prompt (blocking per WS loop — one prompt at a time)
195-
handlePrompt(ctx, conn, store, resources, resolved, system, msg.Content, msg.SessionID)
209+
// Handle session switch mid-connection (new conversation)
210+
if msg.SessionID != "" && (currentSession == nil || currentSession.ID != msg.SessionID) {
211+
sess, err := store.Load(msg.SessionID)
212+
if err == nil {
213+
currentSession = sess
214+
// Restore buffer from the resumed session
215+
if mm := agent.Memory(); mm != nil && len(sess.Buffer) > 0 {
216+
mm.RestoreBuffer(sess.Buffer)
217+
}
218+
}
219+
}
220+
221+
// Run prompt — passes the persistent agent for buffer continuity
222+
currentSession = handlePrompt(ctx, conn, store, resources, resolved, agent, currentSession, msg.Content, msg.SessionID)
223+
}
224+
225+
// WebSocket disconnected — extract episode if enough turns
226+
if currentSession != nil {
227+
if mm := agent.Memory(); mm != nil {
228+
msgStrs := makeSessionMessageStrings(currentSession)
229+
mm.OnSessionEnd(currentSession.ID, currentSession.Turns, msgStrs)
230+
}
196231
}
197232
}
198233
}
199234

200-
func handlePrompt(ctx context.Context, conn *ws.Conn, store *session.Store, resources *resource.Registry, resolved config.ResolvedConfig, system string, prompt string, sessionID string) {
235+
// ── Prompt Handler ─────────────────────────────────────────────────────
236+
237+
// handlePrompt processes a single user prompt within a WebSocket connection.
238+
// Uses the persistent agent (for buffer continuity) and manages session state.
239+
// Returns the updated session (may be a new session for first prompts).
240+
func handlePrompt(
241+
ctx context.Context,
242+
conn *ws.Conn,
243+
store *session.Store,
244+
resources *resource.Registry,
245+
resolved config.ResolvedConfig,
246+
agent *kode.Agent,
247+
currSess *session.Session,
248+
prompt string,
249+
sessionID string,
250+
) *session.Session {
201251
// Resolve @ references
202252
refs := resource.ParseRefs(prompt)
203253
resolvedRefs := make(map[string]string)
@@ -213,43 +263,31 @@ func handlePrompt(ctx context.Context, conn *ws.Conn, store *session.Store, reso
213263
// Load or create session
214264
var sess *session.Session
215265
var err error
266+
216267
if sessionID != "" {
217268
sess, err = store.Load(sessionID)
218269
if err != nil {
219270
sess = nil
220271
}
221272
}
222273

223-
// Build agent for this prompt
224-
agent, sandboxCleanup, err := newServeAgent(resolved, system)
225-
if err != nil {
226-
writeWSError(conn, fmt.Sprintf("agent: %v", err))
227-
return
228-
}
229-
defer agent.Close()
230-
if sandboxCleanup != nil {
231-
defer sandboxCleanup()
232-
}
233-
234274
// Build message history
235275
var messages []llm.Message
276+
isNewSession := false
277+
236278
if sess != nil {
237279
messages = sess.GetMessages()
238280
messages = append(messages, llm.Message{Role: "user", Content: enrichedPrompt})
239281
} else {
240-
modelLabel := kode.ProfileLabel(resolved.Model)
241-
if modelLabel == "" {
242-
modelLabel = "kode"
243-
}
244-
282+
isNewSession = true
245283
messages = []llm.Message{
246-
{Role: "system", Content: system},
284+
{Role: "system", Content: ""},
247285
{Role: "user", Content: enrichedPrompt},
248286
}
249287

250288
// Persist new session
251289
newSess, err := store.Create(
252-
[]llm.Message{{Role: "system", Content: system}},
290+
[]llm.Message{{Role: "system", Content: ""}},
253291
resolved.Model,
254292
shorten(prompt, 60),
255293
)
@@ -267,24 +305,21 @@ func handlePrompt(ctx context.Context, conn *ws.Conn, store *session.Store, reso
267305
}
268306
writeWSJSON(conn, map[string]any{"type": "session", "session_id": sid, "model": resolved.Model})
269307

270-
// Run agent with WebSocket streaming writer
271-
rend := render.New(&wsStreamWriter{conn: conn}, false)
272-
_ = rend // agent doesn't expose renderer injection post-creation
308+
// Append user input to buffer
309+
if mm := agent.Memory(); mm != nil {
310+
userSummary := shorten(prompt, 100)
311+
mm.AppendBuffer("user", userSummary)
312+
}
273313

274-
// The agent loop runs silently (no renderer). We stream the final
275-
// answer via the wsStreamWriter as we get it.
314+
// Run agent
276315
origLen := len(messages) - 1 // exclude the user message we just appended
277-
278-
// Since the agent uses RunWithMessages (not streaming per-token),
279-
// we send the final result as a single "token" event.
280-
// TODO: add streaming callback support to the agent loop
281316
start := time.Now()
282317
_, allMessages, err := agent.RunWithMessages(ctx, messages)
283318
latency := time.Since(start)
284319

285320
if err != nil {
286321
writeWSError(conn, err.Error())
287-
return
322+
return currSess // return unchanged session on error
288323
}
289324

290325
// Get new messages from this turn
@@ -321,15 +356,37 @@ func handlePrompt(ctx context.Context, conn *ws.Conn, store *session.Store, reso
321356
}
322357
}
323358

359+
// Find the assistant response for buffer
360+
if mm := agent.Memory(); mm != nil {
361+
for _, msg := range newMsgs {
362+
if msg.Role == "assistant" && msg.Content != "" {
363+
agentSummary := shorten(msg.Content, 100)
364+
mm.AppendBuffer("agent", agentSummary)
365+
break
366+
}
367+
}
368+
}
369+
324370
writeWSJSON(conn, map[string]any{
325371
"type": "done",
326372
"latency": latency.Seconds(),
327373
})
328374

329-
// Save session — append new messages
375+
// Save session — persist messages AND buffer
330376
if sess != nil {
377+
// Persist buffer to session before saving
378+
if mm := agent.Memory(); mm != nil {
379+
sess.Buffer = mm.GetBuffer()
380+
}
331381
store.Append(sess.ID, newMsgs)
332382
}
383+
384+
// If we started a new session, return it so the WebSocket loop
385+
// tracks it for future turns and OnSessionEnd.
386+
if isNewSession && sess != nil {
387+
return sess
388+
}
389+
return currSess
333390
}
334391

335392
// ── WebSocket Stream Writer ─────────────────────────────────────────────
@@ -421,6 +478,15 @@ func handleStatic() http.HandlerFunc {
421478

422479
// ── Helpers ────────────────────────────────────────────────────────────
423480

481+
func makeSessionMessageStrings(sess *session.Session) []string {
482+
msgs := sess.GetMessages()
483+
out := make([]string, 0, len(msgs))
484+
for _, m := range msgs {
485+
out = append(out, m.Role+": "+m.Content)
486+
}
487+
return out
488+
}
489+
424490
func openInBrowser(url string) {
425491
cmds := []string{"xdg-open", "open", "gnome-open"}
426492
for _, cmd := range cmds {

kode_test.go

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -337,11 +337,16 @@ func TestNew_WithTools(t *testing.T) {
337337
if len(tools) != 2 {
338338
t.Fatalf("expected 2 tools (test_tool + memory) in registry, got %d", len(tools))
339339
}
340-
if tools[0].Name() != "test_tool" {
341-
t.Errorf("tool[0] name = %q, want %q", tools[0].Name(), "test_tool")
340+
// Map iteration is non-deterministic, so check by name
341+
names := make(map[string]bool)
342+
for _, t := range tools {
343+
names[t.Name()] = true
342344
}
343-
if tools[1].Name() != "memory" {
344-
t.Errorf("tool[1] name = %q, want %q", tools[1].Name(), "memory")
345+
if !names["test_tool"] {
346+
t.Errorf("expected test_tool in registry, got %v", names)
347+
}
348+
if !names["memory"] {
349+
t.Errorf("expected memory tool in registry, got %v", names)
345350
}
346351
}
347352

0 commit comments

Comments
 (0)