Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 6 additions & 5 deletions internal/domain/preview.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,11 +9,12 @@ const (
// frames are NEVER persisted and never appear in the event history; they are
// delivered only to stream subscribers that opted in via event_deltas[].
type PreviewFrame struct {
Kind string // PreviewEventStart | PreviewEventDelta
EventID string // the id of the event being previewed (== the eventual persisted event id)
EventType string // event_start: the previewed event's type (e.g. "agent.message")
Index int // event_delta: content block index
Text string // event_delta: incremental text
Kind string // PreviewEventStart | PreviewEventDelta
EventID string // the id of the event being previewed (== the eventual persisted event id)
EventType string // event_start: the previewed event's type (e.g. "agent.message")
ModelRequestStartID string // internal correlation fence; deliberately omitted from WireJSON
Index int // event_delta: content block index
Text string // event_delta: incremental text
}

func (f PreviewFrame) WireJSON() map[string]any {
Expand Down
10 changes: 8 additions & 2 deletions internal/domain/preview_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,10 @@ import (
)

func TestPreviewFrame_WireJSON_Start(t *testing.T) {
f := PreviewFrame{Kind: PreviewEventStart, EventID: "sevt_1", EventType: "agent.message"}
f := PreviewFrame{
Kind: PreviewEventStart, EventID: "sevt_1", EventType: "agent.message",
ModelRequestStartID: "sevt_model_start",
}
got := f.WireJSON()
want := map[string]any{
"type": "event_start",
Expand All @@ -18,7 +21,10 @@ func TestPreviewFrame_WireJSON_Start(t *testing.T) {
}

func TestPreviewFrame_WireJSON_Delta(t *testing.T) {
f := PreviewFrame{Kind: PreviewEventDelta, EventID: "sevt_1", Index: 0, Text: "Hi"}
f := PreviewFrame{
Kind: PreviewEventDelta, EventID: "sevt_1", Index: 0, Text: "Hi",
ModelRequestStartID: "sevt_model_start",
}
got := f.WireJSON()
want := map[string]any{
"type": "event_delta",
Expand Down
34 changes: 32 additions & 2 deletions internal/live/integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -155,7 +155,8 @@ func TestNATSStreamReconcilesLedgerAndCarriesPreviews(t *testing.T) {
}
orderedPreview := domain.PreviewFrame{
Kind: domain.PreviewEventDelta, EventID: "sevt_live_preview_order",
EventType: domain.EvAgentMessage, Index: 0, Text: "first delta",
EventType: domain.EvAgentMessage, ModelRequestStartID: start.ID,
Index: 0, Text: "first delta",
}
if err := broker.PublishPreview(ctx, orderedSession.ID, orderedPreview); err != nil {
t.Fatal(err)
Expand Down Expand Up @@ -188,13 +189,42 @@ func TestNATSStreamReconcilesLedgerAndCarriesPreviews(t *testing.T) {
if ordered.Event == nil || ordered.Event.ID != end.ID {
t.Fatalf("preview-closing frame = %+v, want end event %s", ordered, end.ID)
}
nextStart := domain.EventDraft{
ID: "sevt_live_next_model_start", Type: domain.EvSpanModelRequestStart,
Payload: map[string]any{},
}
if err := store.AppendWorkflowEvents(
ctx,
orderedSession.ID,
orderedAdmission.Events[0].ID,
[]domain.EventDraft{nextStart},
); err != nil {
t.Fatal(err)
}
ordered = receiveFrame(t, orderedFrames)
if ordered.Event == nil || ordered.Event.ID != nextStart.ID {
t.Fatalf("next model start frame = %+v, want %s", ordered, nextStart.ID)
}
if err := broker.PublishPreview(ctx, orderedSession.ID, domain.PreviewFrame{
Kind: domain.PreviewEventDelta, EventID: orderedPreview.EventID,
EventType: domain.EvAgentMessage, Index: 0, Text: "late delta",
EventType: domain.EvAgentMessage, ModelRequestStartID: start.ID,
Index: 0, Text: "late delta",
}); err != nil {
t.Fatal(err)
}
assertNoFrame(t, orderedFrames)
activePreview := domain.PreviewFrame{
Kind: domain.PreviewEventDelta, EventID: "sevt_live_next_preview",
EventType: domain.EvAgentMessage, ModelRequestStartID: nextStart.ID,
Index: 0, Text: "active delta",
}
if err := broker.PublishPreview(ctx, orderedSession.ID, activePreview); err != nil {
t.Fatal(err)
}
ordered = receiveFrame(t, orderedFrames)
if ordered.Preview == nil || ordered.Preview.Text != activePreview.Text {
t.Fatalf("active model preview = %+v, want %+v", ordered, activePreview)
}
cancelOrdered()

// Core NATS is at-most-once. Drop the publisher deliberately and prove the
Expand Down
43 changes: 29 additions & 14 deletions internal/live/nats.go
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,9 @@ func (b *Broker) PublishPreview(
) error {
payload, err := json.Marshal(previewEnvelope{
Kind: frame.Kind, EventID: frame.EventID, EventType: frame.EventType,
Index: frame.Index, Text: frame.Text,
ModelRequestStartID: frame.ModelRequestStartID,
Index: frame.Index,
Text: frame.Text,
})
if err != nil {
return err
Expand All @@ -74,11 +76,12 @@ func (b *Broker) PublishPreview(
}

type previewEnvelope struct {
Kind string `json:"kind"`
EventID string `json:"event_id"`
EventType string `json:"event_type"`
Index int `json:"index,omitempty"`
Text string `json:"text,omitempty"`
Kind string `json:"kind"`
EventID string `json:"event_id"`
EventType string `json:"event_type"`
ModelRequestStartID string `json:"model_request_start_id,omitempty"`
Index int `json:"index,omitempty"`
Text string `json:"text,omitempty"`
}

func eventSubject(sessionID string) string {
Expand Down Expand Up @@ -201,7 +204,8 @@ func (s *Stream) tail(
defer ticker.Stop()
previewPrepared := make(map[string]struct{})
closedPreviewEvents := make(map[string]struct{})
modelRequestClosed := false
closedModelRequests := make(map[string]struct{})
legacyModelRequestClosed := false

reconcile := func() bool {
for {
Expand All @@ -213,9 +217,12 @@ func (s *Stream) tail(
cursor = event.Sequence
switch event.Type {
case domain.EvSpanModelRequestStart:
modelRequestClosed = false
legacyModelRequestClosed = false
case domain.EvSpanModelRequestEnd:
modelRequestClosed = true
legacyModelRequestClosed = true
if startID, _ := event.Payload["model_request_start_id"].(string); startID != "" {
closedModelRequests[startID] = struct{}{}
}
case domain.EvAgentMessage:
closedPreviewEvents[event.ID] = struct{}{}
}
Expand Down Expand Up @@ -289,16 +296,24 @@ func (s *Stream) tail(
// buffered agent.message.
continue
}
if modelRequestClosed {
if envelope.ModelRequestStartID != "" {
if _, closed := closedModelRequests[envelope.ModelRequestStartID]; closed {
// A later model request must not reopen frames buffered for an
// earlier failed or interrupted request.
continue
}
} else if legacyModelRequestClosed {
// Error and interrupt paths intentionally have no authoritative
// agent.message. Once their durable span end has reached this
// subscriber, it still closes any preview frames buffered on the
// separate NATS subject.
// agent.message. Retain the previous global fence for previews from
// older publishers that do not carry request correlation yet.
continue
}
preview := domain.PreviewFrame{
Kind: envelope.Kind, EventID: envelope.EventID,
EventType: envelope.EventType, Index: envelope.Index, Text: envelope.Text,
EventType: envelope.EventType,
ModelRequestStartID: envelope.ModelRequestStartID,
Index: envelope.Index,
Text: envelope.Text,
}
select {
case frames <- app.Frame{Preview: &preview}:
Expand Down
18 changes: 10 additions & 8 deletions internal/temporal/activities.go
Original file line number Diff line number Diff line change
Expand Up @@ -1186,18 +1186,20 @@ func (a *Activities) CallModel(ctx context.Context, in CallModelInput) (CallMode
defer previewMu.Unlock()
if !startedPreview {
_ = a.previews.PublishPreview(ctx, in.SessionID, domain.PreviewFrame{
Kind: domain.PreviewEventStart,
EventID: messageEventID,
EventType: domain.EvAgentMessage,
Kind: domain.PreviewEventStart,
EventID: messageEventID,
EventType: domain.EvAgentMessage,
ModelRequestStartID: modelRequestStartID,
})
startedPreview = true
}
_ = a.previews.PublishPreview(ctx, in.SessionID, domain.PreviewFrame{
Kind: domain.PreviewEventDelta,
EventID: messageEventID,
EventType: domain.EvAgentMessage,
Index: index,
Text: text,
Kind: domain.PreviewEventDelta,
EventID: messageEventID,
EventType: domain.EvAgentMessage,
ModelRequestStartID: modelRequestStartID,
Index: index,
Text: text,
})
})
if err != nil {
Expand Down
10 changes: 9 additions & 1 deletion internal/temporal/preview_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,8 @@ func TestCallModelPublishesCorrelatedPreviewFrames(t *testing.T) {
}
if frames[0].sessionID != "sesn_1" ||
frames[0].frame.Kind != domain.PreviewEventStart ||
frames[0].frame.EventID != result.MessageEventID {
frames[0].frame.EventID != result.MessageEventID ||
frames[0].frame.ModelRequestStartID != result.ModelRequestStartID {
t.Fatalf("start frame = %+v, result id=%s", frames[0], result.MessageEventID)
}
var text strings.Builder
Expand All @@ -51,6 +52,13 @@ func TestCallModelPublishesCorrelatedPreviewFrames(t *testing.T) {
if frame.frame.EventID != result.MessageEventID {
t.Fatalf("delta id = %s, want %s", frame.frame.EventID, result.MessageEventID)
}
if frame.frame.ModelRequestStartID != result.ModelRequestStartID {
t.Fatalf(
"delta model request id = %s, want %s",
frame.frame.ModelRequestStartID,
result.ModelRequestStartID,
)
}
text.WriteString(frame.frame.Text)
}
if got, want := text.String(), "echo: hello"; got != want {
Expand Down