diff --git a/go/adk/pkg/README.md b/go/adk/pkg/README.md index b5737f8e55..8e01e0a368 100644 --- a/go/adk/pkg/README.md +++ b/go/adk/pkg/README.md @@ -14,19 +14,20 @@ Shared types, interfaces, and implementations for the Kagent Go ADK. - **runner/** - Google ADK `runner.Config` creation from `AgentConfig` - **session/** - Session management, persistence, and ADK session service adapter - **skills/** - Agent skills discovery and shell execution -- **taskstore/** - Task storage and A2A result aggregation +- **taskstore/** - A2A task persistence through the kagent controller API - **telemetry/** - OpenTelemetry tracing utilities ## Event Processing -The executor (`KAgentExecutor`) holds a `*runner.Runner` directly and implements `a2asrv.AgentExecutor`: +The executor (`KAgentExecutor`) is a thin kagent-specific wrapper around the +upstream `adka2a.Executor`: ``` -main.go -> CreateGoogleADKRunner -> *runner.Runner +main.go -> CreateRunnerConfig -> runner.Config | -KAgentExecutor.Execute(ctx, reqCtx, queue) - -> runner.Run(ctx, userID, sessionID, content, runConfig) - -> iterate *adksession.Event - -> ConvertADKEventToA2AEvents -> queue.Write - -> inline aggregation -> final status/artifact +KAgentExecutor.Execute(ctx, reqCtx) + -> kagent auth, telemetry, skills, session state, HITL resume setup + -> adka2a.Executor.Execute(ctx, reqCtx) + -> artifact updates for task output + -> status-only lifecycle and terminal events ``` diff --git a/go/adk/pkg/a2a/converter.go b/go/adk/pkg/a2a/converter.go index 616fa1109e..4aca1972bb 100644 --- a/go/adk/pkg/a2a/converter.go +++ b/go/adk/pkg/a2a/converter.go @@ -2,8 +2,6 @@ package a2a import ( "context" - "encoding/json" - "maps" a2atype "github.com/a2aproject/a2a-go/v2/a2a" "google.golang.org/adk/v2/server/adka2a/v2" @@ -19,42 +17,6 @@ func isEmptyDataPart(part *a2atype.Part) bool { return dp != nil && len(dp) == 0 } -// filterTextParts returns only TextParts from the given parts. -func filterTextParts(parts a2atype.ContentParts) a2atype.ContentParts { - var out a2atype.ContentParts - for _, p := range parts { - if p != nil && p.Text() != "" { - out = append(out, p) - } - } - return out -} - -// messageToGenAIContent converts an A2A message to *genai.Content using kagent -// a2aPartConverter logic: handle kagent_type and adk_type DataParts explicitly, -// drop unrecognised DataParts (e.g. HITL decision parts). -func messageToGenAIContent(ctx context.Context, msg *a2atype.Message) (*genai.Content, error) { - if msg == nil { - return nil, nil - } - parts := make([]*genai.Part, 0, len(msg.Parts)) - for _, part := range msg.Parts { - genaiPart, err := a2aPartConverter(ctx, msg, part) - if err != nil { - return nil, err - } - if genaiPart == nil { - continue - } - parts = append(parts, genaiPart) - } - var role genai.Role = genai.RoleUser - if msg.Role == a2atype.MessageRoleAgent { - role = genai.RoleModel - } - return genai.NewContentFromParts(parts, role), nil -} - // a2aPartConverter converts inbound A2A parts to GenAI parts. func a2aPartConverter(_ context.Context, _ a2atype.Event, part *a2atype.Part) (*genai.Part, error) { dp := asDataPart(part) @@ -82,6 +44,19 @@ func a2aPartConverter(_ context.Context, _ a2atype.Event, part *a2atype.Part) (* return nil, nil } +// genAIPartConverter lets the upstream executor own artifact construction +// while preserving kagent's part filtering and long-running-tool metadata. +func genAIPartConverter(_ context.Context, event *adksession.Event, part *genai.Part) (*a2atype.Part, error) { + converted, err := adka2a.ToA2APart(part, event.LongRunningToolIDs) + if err != nil { + return nil, err + } + if isEmptyDataPart(converted) { + return nil, nil + } + return converted, nil +} + // convertDataPartToGenAI converts a DataPart with a type metadata key // (either adk_type or kagent_type) back to GenAI for inbound message processing. func convertDataPartToGenAI(data map[string]any, metadata map[string]any, typeKey string) (*genai.Part, error) { @@ -113,46 +88,3 @@ func convertDataPartToGenAI(data map[string]any, metadata map[string]any, typeKe } return adka2a.ToGenAIPart(a2atype.NewDataPart(data)) } - -// toA2AMetadataMap converts v to map[string]any via JSON so values placed in A2A -func toA2AMetadataMap(v any) (map[string]any, error) { - if v == nil { - return nil, nil - } - b, err := json.Marshal(v) - if err != nil { - return nil, err - } - var m map[string]any - if err := json.Unmarshal(b, &m); err != nil { - return nil, err - } - return m, nil -} - -// buildEventMeta merges the base metadata with per-event fields such as -// invocation_id, author, branch, usage_metadata, etc. -func buildEventMeta(baseMeta map[string]any, adkEvent *adksession.Event) map[string]any { - result := maps.Clone(baseMeta) - if adkEvent == nil { - return result - } - for k, v := range map[string]string{ - "invocation_id": adkEvent.InvocationID, - "author": adkEvent.Author, - "branch": adkEvent.Branch, - } { - if v != "" { - result[adka2a.ToA2AMetaKey(k)] = v - } - } - if adkEvent.UsageMetadata != nil { - if um, err := toA2AMetadataMap(adkEvent.UsageMetadata); err == nil && um != nil { - result[adka2a.ToA2AMetaKey("usage_metadata")] = um - } - } - if adkEvent.ErrorCode != "" { - result[adka2a.ToA2AMetaKey("error_code")] = adkEvent.ErrorCode - } - return result -} diff --git a/go/adk/pkg/a2a/converter_test.go b/go/adk/pkg/a2a/converter_test.go index c9c38b1312..88f3ffeffa 100644 --- a/go/adk/pkg/a2a/converter_test.go +++ b/go/adk/pkg/a2a/converter_test.go @@ -6,6 +6,7 @@ import ( a2atype "github.com/a2aproject/a2a-go/v2/a2a" "google.golang.org/adk/v2/server/adka2a/v2" + adksession "google.golang.org/adk/v2/session" "google.golang.org/genai" ) @@ -118,48 +119,35 @@ func TestConvertDataPartToGenAI_UnknownType(t *testing.T) { } // --------------------------------------------------------------------------- -// messageToGenAIContent +// a2aPartConverter // --------------------------------------------------------------------------- -func TestMessageToGenAIContent_TextPart(t *testing.T) { - msg := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("hello")) - content, err := messageToGenAIContent(context.Background(), msg) +func TestA2APartConverter_TextPart(t *testing.T) { + part, err := a2aPartConverter(context.Background(), nil, a2atype.NewTextPart("hello")) if err != nil { t.Fatalf("unexpected error: %v", err) } - if content == nil { - t.Fatal("expected non-nil content") - return - } - if len(content.Parts) != 1 { - t.Fatalf("expected 1 part, got %d", len(content.Parts)) - } - if content.Parts[0].Text != "hello" { - t.Errorf("text = %q, want %q", content.Parts[0].Text, "hello") + if part == nil || part.Text != "hello" { + t.Fatalf("converted part = %#v, want text hello", part) } } -func TestMessageToGenAIContent_DropsUnrecognisedDataPart(t *testing.T) { +func TestA2APartConverter_DropsUnrecognisedDataPart(t *testing.T) { // A DataPart with no recognised kagent_type metadata (e.g. a HITL decision // payload like {decision_type: "approve"}) should be dropped silently. - msg := a2atype.NewMessage(a2atype.MessageRoleUser, - a2atype.NewTextPart("approving"), + part, err := a2aPartConverter( + context.Background(), nil, convDataPart(map[string]any{"decision_type": "approve"}, nil), ) - content, err := messageToGenAIContent(context.Background(), msg) if err != nil { t.Fatalf("unexpected error: %v", err) } - // Only the TextPart should survive; the unrecognised DataPart is dropped. - if len(content.Parts) != 1 { - t.Fatalf("expected 1 part (DataPart dropped), got %d", len(content.Parts)) - } - if content.Parts[0].Text != "approving" { - t.Errorf("remaining part text = %q, want %q", content.Parts[0].Text, "approving") + if part != nil { + t.Fatalf("converted part = %#v, want nil", part) } } -func TestMessageToGenAIContent_KagentTypeFunctionResponse(t *testing.T) { +func TestA2APartConverter_KagentTypeFunctionResponse(t *testing.T) { // A DataPart with kagent_type=function_response should be converted to GenAI. dp := convDataPart(map[string]any{ "name": "my_func", @@ -168,66 +156,33 @@ func TestMessageToGenAIContent_KagentTypeFunctionResponse(t *testing.T) { }, map[string]any{ GetKAgentMetadataKey(A2ADataPartMetadataTypeKey): A2ADataPartMetadataTypeFunctionResponse, }) - msg := a2atype.NewMessage(a2atype.MessageRoleUser, dp) - content, err := messageToGenAIContent(context.Background(), msg) + part, err := a2aPartConverter(context.Background(), nil, dp) if err != nil { t.Fatalf("unexpected error: %v", err) } - if len(content.Parts) != 1 { - t.Fatalf("expected 1 part, got %d", len(content.Parts)) - } - if content.Parts[0].FunctionResponse == nil { + if part == nil || part.FunctionResponse == nil { t.Fatal("expected FunctionResponse, got nil") } - if content.Parts[0].FunctionResponse.Name != "my_func" { - t.Errorf("name = %q, want my_func", content.Parts[0].FunctionResponse.Name) - } -} - -func TestMessageToGenAIContent_NilMessage(t *testing.T) { - content, err := messageToGenAIContent(context.Background(), nil) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if content != nil { - t.Errorf("expected nil content for nil message, got %v", content) + if part.FunctionResponse.Name != "my_func" { + t.Errorf("name = %q, want my_func", part.FunctionResponse.Name) } } -// --------------------------------------------------------------------------- -// toA2AMetadataMap -// --------------------------------------------------------------------------- - -func TestToA2AMetadataMap(t *testing.T) { - t.Parallel() - um := &genai.GenerateContentResponseUsageMetadata{ - PromptTokenCount: 10, - CandidatesTokenCount: 20, - } - m, err := toA2AMetadataMap(um) +func TestGenAIPartConverter_PreservesLongRunningMetadata(t *testing.T) { + call := genai.NewPartFromFunctionCall("dangerous_tool", map[string]any{"path": "/tmp/x"}) + call.FunctionCall.ID = "call-1" + part, err := genAIPartConverter( + context.Background(), + &adksession.Event{LongRunningToolIDs: []string{"call-1"}}, + call, + ) if err != nil { - t.Fatalf("toA2AMetadataMap: %v", err) - } - if m == nil { - t.Fatal("expected non-nil map") + t.Fatalf("genAIPartConverter() error = %v", err) } - pt, ok := m["promptTokenCount"].(float64) - if !ok || pt != 10 { - t.Fatalf("promptTokenCount: got %v (%T), want float64 10", m["promptTokenCount"], m["promptTokenCount"]) - } - ct, ok := m["candidatesTokenCount"].(float64) - if !ok || ct != 20 { - t.Fatalf("candidatesTokenCount: got %v (%T), want float64 20", m["candidatesTokenCount"], m["candidatesTokenCount"]) - } -} - -func TestToA2AMetadataMap_nil(t *testing.T) { - t.Parallel() - m, err := toA2AMetadataMap(nil) - if err != nil { - t.Fatalf("toA2AMetadataMap(nil): %v", err) + if part == nil { + t.Fatal("genAIPartConverter() returned nil") } - if m != nil { - t.Fatalf("expected nil map, got %#v", m) + if got, _ := ReadMetadataValue(part.Metadata, A2ADataPartMetadataIsLongRunningKey); got != true { + t.Fatalf("long-running metadata = %#v, want true", got) } } diff --git a/go/adk/pkg/a2a/executor.go b/go/adk/pkg/a2a/executor.go index 39ad1e9b30..e731b238f5 100644 --- a/go/adk/pkg/a2a/executor.go +++ b/go/adk/pkg/a2a/executor.go @@ -4,7 +4,6 @@ import ( "context" "fmt" "iter" - "maps" "os" "strings" @@ -16,6 +15,7 @@ import ( "github.com/kagent-dev/kagent/go/adk/pkg/skills" "github.com/kagent-dev/kagent/go/adk/pkg/telemetry" "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/trace" adkagent "google.golang.org/adk/v2/agent" "google.golang.org/adk/v2/runner" "google.golang.org/adk/v2/server/adka2a/v2" @@ -28,7 +28,7 @@ const ( sessionNameMaxLength = 20 ) -// KAgentExecutorConfig holds the configuration for KAgentExecutor +// KAgentExecutorConfig holds the configuration for KAgentExecutor. type KAgentExecutorConfig struct { RunnerConfig runner.Config SessionService adksession.Service @@ -38,11 +38,11 @@ type KAgentExecutorConfig struct { Logger logr.Logger } -// KAgentExecutor implements a2asrv.AgentExecutor +// KAgentExecutor keeps kagent's request/session glue around the upstream ADK +// A2A executor. Event conversion and artifact streaming are delegated to ADK. type KAgentExecutor struct { - runnerConfig runner.Config + builtin a2asrv.AgentExecutor sessionService adksession.Service - stream bool appName string skillsDirectory string logger logr.Logger @@ -50,7 +50,7 @@ type KAgentExecutor struct { var _ a2asrv.AgentExecutor = (*KAgentExecutor)(nil) -// NewKAgentExecutor creates a KAgentExecutor from config +// NewKAgentExecutor creates a KAgentExecutor from config. func NewKAgentExecutor(cfg KAgentExecutorConfig) *KAgentExecutor { skillsDir := cfg.SkillsDirectory if skillsDir == "" { @@ -59,10 +59,32 @@ func NewKAgentExecutor(cfg KAgentExecutorConfig) *KAgentExecutor { if skillsDir == "" { skillsDir = defaultSkillsDirectory } + + var runConfig adkagent.RunConfig + if cfg.Stream { + runConfig.StreamingMode = adkagent.StreamingModeSSE + } + runnerConfig := cfg.RunnerConfig + if cfg.SessionService != nil { + runnerConfig.SessionService = cfg.SessionService + } + builtin := adka2a.NewExecutor(adka2a.ExecutorConfig{ + RunnerConfig: runnerConfig, + RunConfig: runConfig, + A2APartConverter: a2aPartConverter, + GenAIPartConverter: genAIPartConverter, + AfterEventCallback: func(ctx adka2a.ExecutorContext, event *adksession.Event, _ *a2atype.TaskArtifactUpdateEvent) error { + if event.InvocationID != "" { + trace.SpanFromContext(ctx).SetAttributes(attribute.String("gcp.vertex.agent.invocation_id", event.InvocationID)) + } + return nil + }, + OutputMode: adka2a.OutputArtifactPerEvent, + }) + return &KAgentExecutor{ - runnerConfig: cfg.RunnerConfig, - sessionService: cfg.SessionService, - stream: cfg.Stream, + builtin: builtin, + sessionService: runnerConfig.SessionService, appName: cfg.AppName, skillsDirectory: skillsDir, logger: cfg.Logger.WithName("kagent-executor"), @@ -97,36 +119,8 @@ func (u *userIDInterceptor) Before(ctx context.Context, callCtx *a2asrv.CallCont return ctx, nil, nil } -// newAgentMessage builds an agent message stamped with the request's context -// and task ids. A2A allows omitting these (the task is the canonical carrier), -// but stamping them lets consumers that flatten task.history into standalone -// messages key each message to its task without backfilling. Mirrors the Python -// kagent-adk event converter. -func newAgentMessage(reqCtx *a2asrv.ExecutorContext, parts ...*a2atype.Part) *a2atype.Message { - msg := a2atype.NewMessage(a2atype.MessageRoleAgent, parts...) - msg.ContextID = reqCtx.ContextID - msg.TaskID = reqCtx.TaskID - return msg -} - -// newAgentStatusEvent builds a working TaskStatusUpdateEvent whose agent message -// carries the given parts, the given metadata, and the request's context/task -// ids (via newAgentMessage). The message and event share the same metadata map, -// matching the executor's emission paths. This is the per-event seam where a -// streamed agent message is turned into an emitted (and persisted) event, so the -// id stamping here is what the send guard relies on when it later flattens -// task.history. -func newAgentStatusEvent(reqCtx *a2asrv.ExecutorContext, parts a2atype.ContentParts, meta map[string]any) *a2atype.TaskStatusUpdateEvent { - msg := newAgentMessage(reqCtx, parts...) - msg.Metadata = meta - statusEv := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateWorking, msg) - statusEv.Metadata = meta - return statusEv -} - -// Execute implements a2asrv.AgentExecutor. -// It follows the Python _handle_request pattern: set up session, handle HITL, -// convert inbound message, run the agent loop, and emit A2A events. +// Execute applies kagent-specific request setup and delegates event generation +// to the upstream ADK executor, which streams output as artifact updates. func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.ExecutorContext) iter.Seq2[a2atype.Event, error] { return func(yield func(a2atype.Event, error) bool) { if reqCtx.Message == nil { @@ -134,26 +128,14 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.ExecutorCon return } - // 1. Derive userID / sessionID. userID := "A2A_USER_" + reqCtx.ContextID - if callCtx, ok := a2asrv.CallContextFrom(ctx); ok { - if callCtx.User != nil && callCtx.User.Name != "" { - userID = callCtx.User.Name - } + if callCtx, ok := a2asrv.CallContextFrom(ctx); ok && callCtx.User != nil && callCtx.User.Name != "" { + userID = callCtx.User.Name } sessionID := reqCtx.ContextID ctx = withBearerToken(ctx) ctx = auth.WithUserID(ctx, userID) - - e.logger.Info("Execute", - "taskID", reqCtx.TaskID, - "contextID", reqCtx.ContextID, - "appName", e.appName, - "userID", userID, - ) - - // 2. Set up telemetry span attributes. spanAttributes := map[string]string{ "kagent.user_id": userID, "gen_ai.task.id": string(reqCtx.TaskID), @@ -165,10 +147,15 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.ExecutorCon ctx = telemetry.SetKAgentSpanAttributes(ctx, spanAttributes) ctx, invocationSpan := telemetry.StartInvocationSpan(ctx) defer invocationSpan.End() - telemetry.SetMessageMetadataAttributes(ctx, reqCtx.Message.Metadata) - // 3. Initialize skills session path. + e.logger.Info("Execute", + "taskID", reqCtx.TaskID, + "contextID", reqCtx.ContextID, + "appName", e.appName, + "userID", userID, + ) + if e.skillsDirectory != "" && sessionID != "" { if _, err := skills.InitializeSessionPath(sessionID, e.skillsDirectory); err != nil { e.logger.V(1).Info("Skills session path init failed (continuing)", @@ -176,242 +163,82 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.ExecutorCon } } - // 4. Create / lookup session via sessionService. - if e.sessionService != nil { - var sess adksession.Session - resp, err := e.sessionService.Get(ctx, &adksession.GetRequest{AppName: e.appName, UserID: userID, SessionID: sessionID}) - if err != nil { - e.logger.V(1).Info("Session lookup failed, will create", "error", err, "sessionID", sessionID) - } else if resp != nil { - sess = resp.Session - } - if sess == nil { - sessionName := extractSessionName(reqCtx.Message) - state := make(map[string]any) - if sessionName != "" { - state[StateKeySessionName] = sessionName - } - // Propagate x-kagent-source so the session is tagged in the DB. - if callCtx, ok := a2asrv.CallContextFrom(ctx); ok { - if meta := callCtx.ServiceParams(); meta != nil { - if vals, ok := meta.Get("x-kagent-source"); ok && len(vals) > 0 && vals[0] != "" { - state[StateKeySource] = vals[0] - } - } - } - if _, err := e.sessionService.Create(ctx, &adksession.CreateRequest{ - AppName: e.appName, - UserID: userID, - State: state, - SessionID: sessionID, - }); err != nil { - yield(nil, fmt.Errorf("failed to create session: %w", err)) - return - } - } - } - - // 5. Detect HITL decision and build the resume message if needed. - inboundMessage := reqCtx.Message - if resumeMessage := BuildResumeHITLMessage(reqCtx.StoredTask, inboundMessage); resumeMessage != nil { - inboundMessage = resumeMessage - } - - // 6. Convert inbound message to *genai.Content using kagent a2aPartConverter. - content, err := messageToGenAIContent(ctx, inboundMessage) - if err != nil { - yield(nil, fmt.Errorf("inbound message conversion failed: %w", err)) + // Run our own session management before upstream executor runs its prepareSession function. + // This ensures that we create a session that contains metadata like x-kagent-source, + // and the upstream executor will find this session already exists and skip creation. + if err := e.ensureSession(ctx, reqCtx.Message, userID, sessionID); err != nil { + yield(nil, err) return } - // 8. Create runner. - r, err := runner.New(e.runnerConfig) - if err != nil { - yield(nil, fmt.Errorf("failed to create runner: %w", err)) - return - } - - // 9. Emit initial events. - if reqCtx.StoredTask == nil { - if !yield(a2atype.NewSubmittedTask(reqCtx, reqCtx.Message), nil) { - return - } - } else if ExtractDecisionFromMessage(reqCtx.Message) != "" { - // a2a-go appends incoming message to task history before executor runs. - // Remove the pre-appended copy and emit one decision status event. + if ExtractDecisionFromMessage(reqCtx.Message) != "" { + // a2a-go appends the inbound decision before invoking the executor. The + // original decision is re-emitted once for history/audit, while the + // transformed FunctionResponses are what ADK must consume. dropPreAppendedDecisionFromHistory(reqCtx.StoredTask, reqCtx.Message) decision := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateWorking, reqCtx.Message) if !yield(decision, nil) { return } - } - - // Base metadata carried on every event (app_name, user_id, session_id). - baseMeta := map[string]any{ - adka2a.ToA2AMetaKey("app_name"): e.appName, - adka2a.ToA2AMetaKey("user_id"): userID, - adka2a.ToA2AMetaKey("session_id"): sessionID, - } - - working := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateWorking, nil) - working.Metadata = maps.Clone(baseMeta) - if !yield(working, nil) { - return - } - - // 10. Run the agent event loop. - var runConfig adkagent.RunConfig - if e.stream { - runConfig.StreamingMode = adkagent.StreamingModeSSE - } - - // State tracked across the event loop. - var ( - invocationID string - lastNonPartialParts a2atype.ContentParts - hitlParts a2atype.ContentParts - runErr error - ) - - for adkEvent, adkErr := range r.Run(ctx, userID, sessionID, content, runConfig) { - if adkErr != nil { - runErr = adkErr - break - } - if adkEvent == nil { - continue - } - - // Track invocation ID from the first event that has one. - if adkEvent.InvocationID != "" && invocationID == "" { - invocationID = adkEvent.InvocationID - invocationSpan.SetAttributes(attribute.String("gcp.vertex.agent.invocation_id", invocationID)) - } - - // Build per-event metadata (inherits baseMeta + adds invocation_id, usage etc.). - eventMeta := buildEventMeta(baseMeta, adkEvent) - - // Convert GenAI parts → A2A parts (with kagent stamping). - if adkEvent.Content == nil || len(adkEvent.Content.Parts) == 0 { - if adkEvent.ErrorCode != "" { - errMsg := newAgentMessage(reqCtx, - a2atype.NewTextPart(fmt.Sprintf("LLM error: %s %s", adkEvent.ErrorCode, adkEvent.ErrorMessage))) - failed := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateFailed, errMsg) - failed.Metadata = eventMeta - yield(failed, nil) - return - } - continue + // Transform the Kagent-specific decision message into a resume message for the upstream executor. + // The ADK HITL resume is handled upstream in the HandleInputRequired function. + if resumeMessage := BuildResumeHITLMessage(reqCtx.StoredTask, reqCtx.Message); resumeMessage != nil { + reqCtx.Message = resumeMessage } + } - if adkEvent.ErrorCode != "" { - errMsg := newAgentMessage(reqCtx, - a2atype.NewTextPart(fmt.Sprintf("LLM error: %s %s", adkEvent.ErrorCode, adkEvent.ErrorMessage))) - failed := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateFailed, errMsg) - failed.Metadata = eventMeta - yield(failed, nil) + for event, err := range e.builtin.Execute(ctx, reqCtx) { + if !yield(event, err) { return } - - var a2aParts a2atype.ContentParts - for _, genaiPart := range adkEvent.Content.Parts { - if genaiPart == nil { - continue - } - a2aPart, err := adka2a.ToA2APart(genaiPart, adkEvent.LongRunningToolIDs) - if err != nil { - continue - } - if isEmptyDataPart(a2aPart) { - continue - } - a2aParts = append(a2aParts, a2aPart) - } - - // Collect HITL (input_required) parts from LongRunningToolIDs. - isHITLEvent := len(adkEvent.LongRunningToolIDs) > 0 - if isHITLEvent { - hitlParts = append(hitlParts, a2aParts...) - } - - if len(a2aParts) == 0 { - continue - } - - if adkEvent.Partial { - // Partial event: emit as working status (text-only) for UI streaming. - // Note: Go ADK executor uses TaskArtifactUpdateEvent for partial events, - // so we don't need to emit a separate partial artifact update. - // However, this is done here in order to match the Python executor's behavior. - // Go ADK executor also uses different A2A response formats than Python ADK. - textOnly := filterTextParts(a2aParts) - if len(textOnly) > 0 { - mirrorMeta := maps.Clone(eventMeta) - mirrorMeta[adka2a.ToA2AMetaKey("partial")] = true - statusEv := newAgentStatusEvent(reqCtx, textOnly, mirrorMeta) - if !yield(statusEv, nil) { - return - } - } - } else { - if len(hitlParts) == 0 { - statusEv := newAgentStatusEvent(reqCtx, a2aParts, maps.Clone(eventMeta)) - if !yield(statusEv, nil) { - return - } - lastNonPartialParts = a2aParts - } - } - - // Break on confirmation events that have long-running tool IDs. - if isHITLEvent { - break - } - } - - // 11. Emit final event. - finalMeta := maps.Clone(baseMeta) - if invocationID != "" { - finalMeta[adka2a.ToA2AMetaKey("invocation_id")] = invocationID - } - - if runErr != nil { - errMsg := newAgentMessage(reqCtx, a2atype.NewTextPart(runErr.Error())) - failed := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateFailed, errMsg) - failed.Metadata = finalMeta - yield(failed, nil) - return } + } +} - if len(hitlParts) > 0 { - // input_required: the agent is waiting for HITL decisions. - hitlMsg := newAgentMessage(reqCtx, hitlParts...) - inputRequired := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateInputRequired, hitlMsg) - inputRequired.Metadata = finalMeta - yield(inputRequired, nil) - return - } +// ensureSession ensures that a session exists for the given user and session ID. +// If a session does not exist, it creates a new session with the given user and session ID. +func (e *KAgentExecutor) ensureSession(ctx context.Context, message *a2atype.Message, userID, sessionID string) error { + if e.sessionService == nil { + return nil + } + resp, err := e.sessionService.Get(ctx, &adksession.GetRequest{ + AppName: e.appName, UserID: userID, SessionID: sessionID, + }) + if err == nil && resp != nil && resp.Session != nil { + return nil + } + if err != nil { + e.logger.V(1).Info("Session lookup failed, will create", "error", err, "sessionID", sessionID) + } - // Final artifact update with lastChunk=true (if we have parts) and final completed status update (no message payload). - if len(lastNonPartialParts) > 0 { - finalArtifact := a2atype.NewArtifactEvent(reqCtx, lastNonPartialParts...) - finalArtifact.LastChunk = true - if !yield(finalArtifact, nil) { - return + state := make(map[string]any) + if sessionName := extractSessionName(message); sessionName != "" { + state[StateKeySessionName] = sessionName + } + if callCtx, ok := a2asrv.CallContextFrom(ctx); ok { + if meta := callCtx.ServiceParams(); meta != nil { + if vals, ok := meta.Get("x-kagent-source"); ok && len(vals) > 0 && vals[0] != "" { + state[StateKeySource] = vals[0] } } - - completed := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateCompleted, nil) - completed.Metadata = finalMeta - yield(completed, nil) } + if _, err := e.sessionService.Create(ctx, &adksession.CreateRequest{ + AppName: e.appName, UserID: userID, State: state, SessionID: sessionID, + }); err != nil { + return fmt.Errorf("failed to create session: %w", err) + } + return nil } -// Cancel implements a2asrv.AgentExecutor. +// Cancel delegates cancellation to the upstream executor. func (e *KAgentExecutor) Cancel(ctx context.Context, reqCtx *a2asrv.ExecutorContext) iter.Seq2[a2atype.Event, error] { - return func(yield func(a2atype.Event, error) bool) { - event := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateCanceled, nil) - yield(event, nil) + return e.builtin.Cancel(ctx, reqCtx) +} + +// Cleanup preserves the upstream executor's subagent cleanup behavior. +func (e *KAgentExecutor) Cleanup(ctx context.Context, reqCtx *a2asrv.ExecutorContext, result a2atype.SendMessageResult, cause error) { + if cleaner, ok := e.builtin.(a2asrv.AgentExecutionCleaner); ok { + cleaner.Cleanup(ctx, reqCtx, result, cause) } } diff --git a/go/adk/pkg/a2a/executor_test.go b/go/adk/pkg/a2a/executor_test.go index 1377cca597..c30e8dd86b 100644 --- a/go/adk/pkg/a2a/executor_test.go +++ b/go/adk/pkg/a2a/executor_test.go @@ -1,67 +1,242 @@ package a2a import ( + "context" + "iter" "testing" a2atype "github.com/a2aproject/a2a-go/v2/a2a" "github.com/a2aproject/a2a-go/v2/a2asrv" + "github.com/go-logr/logr" + adkagent "google.golang.org/adk/v2/agent" + "google.golang.org/adk/v2/model" + "google.golang.org/adk/v2/runner" + adksession "google.golang.org/adk/v2/session" + "google.golang.org/genai" ) -// TestNewAgentMessage_StampsContextAndTaskID verifies agent messages carry the -// request's context and task ids. A2A allows omitting them (the task is the -// canonical carrier), but stamping them lets consumers that flatten task.history -// key each message to its task without backfilling. -func TestNewAgentMessage_StampsContextAndTaskID(t *testing.T) { +type recordingExecutor struct { + message *a2atype.Message + cleanupCalled bool + events []a2atype.Event +} + +func (e *recordingExecutor) Execute(_ context.Context, reqCtx *a2asrv.ExecutorContext) iter.Seq2[a2atype.Event, error] { + return func(yield func(a2atype.Event, error) bool) { + e.message = reqCtx.Message + if e.events != nil { + for _, event := range e.events { + if !yield(event, nil) { + return + } + } + return + } + yield(a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateWorking, nil), nil) + } +} + +func (e *recordingExecutor) Cancel(_ context.Context, reqCtx *a2asrv.ExecutorContext) iter.Seq2[a2atype.Event, error] { + return func(yield func(a2atype.Event, error) bool) { + yield(a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateCanceled, nil), nil) + } +} + +func (e *recordingExecutor) Cleanup(context.Context, *a2asrv.ExecutorContext, a2atype.SendMessageResult, error) { + e.cleanupCalled = true +} + +func TestKAgentExecutor_TransformsHITLDecisionBeforeDelegating(t *testing.T) { + decision := a2atype.NewMessage( + a2atype.MessageRoleUser, + dataPart(map[string]any{KAgentHitlDecisionTypeKey: KAgentHitlDecisionTypeApprove}, nil), + ) + storedTask := &a2atype.Task{ + ID: "task-1", + ContextID: "ctx-1", + Status: a2atype.TaskStatus{ + State: a2atype.TaskStateInputRequired, + Message: a2atype.NewMessage( + a2atype.MessageRoleAgent, + dataPart( + map[string]any{ + "name": "adk_request_confirmation", + "id": "confirm-1", + "args": map[string]any{ + "originalFunctionCall": map[string]any{ + "name": "delete_file", + "args": map[string]any{"path": "/tmp/x"}, + "id": "call-1", + }, + }, + }, + map[string]any{ + "kagent_type": "function_call", + "kagent_is_long_running": true, + }, + ), + ), + }, + History: []*a2atype.Message{decision}, + } reqCtx := &a2asrv.ExecutorContext{ - ContextID: "ctx-xyz", - TaskID: a2atype.TaskID("task-xyz"), + TaskID: "task-1", + ContextID: "ctx-1", + Message: decision, + StoredTask: storedTask, } + builtin := &recordingExecutor{} + executor := &KAgentExecutor{builtin: builtin, logger: logr.Discard()} - msg := newAgentMessage(reqCtx, a2atype.NewTextPart("hello")) + var events []a2atype.Event + for event, err := range executor.Execute(context.Background(), reqCtx) { + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + events = append(events, event) + } - if msg.ContextID != "ctx-xyz" { - t.Errorf("ContextID = %q, want %q", msg.ContextID, "ctx-xyz") + if len(events) != 2 { + t.Fatalf("Execute() emitted %d events, want decision acknowledgement and delegated event", len(events)) + } + decisionAck, ok := events[0].(*a2atype.TaskStatusUpdateEvent) + if !ok || decisionAck.Status.State != a2atype.TaskStateWorking || decisionAck.Status.Message != decision { + t.Fatalf("first event = %#v, want original decision acknowledgement", events[0]) + } + working, ok := events[1].(*a2atype.TaskStatusUpdateEvent) + if !ok || working.Status.State != a2atype.TaskStateWorking || working.Status.Message != nil { + t.Fatalf("delegated event = %#v, want content-free working status", events[1]) + } + if len(storedTask.History) != 0 { + t.Fatalf("stored task history len = %d, want pre-appended decision removed", len(storedTask.History)) + } + if builtin.message == nil || len(builtin.message.Parts) != 1 { + t.Fatalf("delegated message = %#v, want one FunctionResponse", builtin.message) } - if msg.TaskID != a2atype.TaskID("task-xyz") { - t.Errorf("TaskID = %q, want %q", msg.TaskID, a2atype.TaskID("task-xyz")) + part := builtin.message.Parts[0] + if got, _ := ReadMetadataValue(part.Metadata, A2ADataPartMetadataTypeKey); got != A2ADataPartMetadataTypeFunctionResponse { + t.Fatalf("delegated part type = %#v, want function_response", got) } - if msg.Role != a2atype.MessageRoleAgent { - t.Errorf("Role = %q, want %q", msg.Role, a2atype.MessageRoleAgent) + if got := asDataPart(part)[PartKeyID]; got != "confirm-1" { + t.Fatalf("delegated FunctionResponse id = %#v, want confirm-1", got) } } -// TestNewAgentStatusEvent_MessageCarriesIDs verifies the per-event emission seam: -// the working status event the executor writes (and that is persisted into -// task.history) carries an agent message stamped with the request's context/task -// ids. This is the property the send guard depends on — without it the persisted -// message keys differently from its locally-streamed counterpart and falsely -// blocks the next send. Mirrors the Python converter test. -func TestNewAgentStatusEvent_MessageCarriesIDs(t *testing.T) { - reqCtx := &a2asrv.ExecutorContext{ - ContextID: "ctx-xyz", - TaskID: a2atype.TaskID("task-xyz"), +func TestKAgentExecutor_ForwardsCleanup(t *testing.T) { + builtin := &recordingExecutor{} + executor := &KAgentExecutor{builtin: builtin} + executor.Cleanup(context.Background(), &a2asrv.ExecutorContext{}, nil, nil) + if !builtin.cleanupCalled { + t.Fatal("Cleanup() was not forwarded to the upstream executor") } - meta := map[string]any{"k": "v"} +} - ev := newAgentStatusEvent(reqCtx, a2atype.ContentParts{a2atype.NewTextPart("hi")}, meta) +func TestKAgentExecutor_PreservesContentBearingLastChunk(t *testing.T) { + reqCtx := &a2asrv.ExecutorContext{TaskID: "task-1", ContextID: "ctx-1"} + reqCtx.Message = a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("hi")) + final := a2atype.NewArtifactEvent(reqCtx, a2atype.NewTextPart("hello")) + final.LastChunk = true + builtin := &recordingExecutor{events: []a2atype.Event{final}} + executor := &KAgentExecutor{builtin: builtin, logger: logr.Discard()} + + var got []a2atype.Event + for event, err := range executor.Execute(context.Background(), reqCtx) { + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + got = append(got, event) + } - if ev.Status.State != a2atype.TaskStateWorking { - t.Errorf("State = %q, want %q", ev.Status.State, a2atype.TaskStateWorking) + if len(got) != 1 || got[0] != final { + t.Fatalf("Execute() events = %#v, want the original final artifact only", got) } - if ev.Status.Message == nil { - t.Fatal("status message is nil") + update, ok := got[0].(*a2atype.TaskArtifactUpdateEvent) + if !ok || !update.LastChunk || len(update.Artifact.Parts) != 1 || update.Artifact.Parts[0].Text() != "hello" { + t.Fatalf("final artifact = %#v, want content-bearing lastChunk event", got[0]) } - if ev.Status.Message.ContextID != "ctx-xyz" { - t.Errorf("message ContextID = %q, want %q", ev.Status.Message.ContextID, "ctx-xyz") +} + +func TestKAgentExecutor_StreamsArtifactsThroughUpstreamExecutor(t *testing.T) { + const ( + appName = "test-app" + contextID = "context-1" + ) + + agent, err := adkagent.New(adkagent.Config{ + Name: "streaming-agent", + Run: func(ic adkagent.InvocationContext) iter.Seq2[*adksession.Event, error] { + return func(yield func(*adksession.Event, error) bool) { + partial := &adksession.Event{ + Author: ic.Agent().Name(), + InvocationID: ic.InvocationID(), + Branch: ic.Branch(), + LLMResponse: model.LLMResponse{ + Content: genai.NewContentFromText("hel", genai.RoleModel), + Partial: true, + }, + } + if !yield(partial, nil) { + return + } + + final := &adksession.Event{ + Author: ic.Agent().Name(), + InvocationID: ic.InvocationID(), + Branch: ic.Branch(), + LLMResponse: model.LLMResponse{ + Content: genai.NewContentFromText("hello", genai.RoleModel), + }, + } + yield(final, nil) + } + }, + }) + if err != nil { + t.Fatalf("agent.New() error = %v", err) + } + + sessionService := adksession.InMemoryService() + executor := NewKAgentExecutor(KAgentExecutorConfig{ + AppName: appName, + SessionService: sessionService, + Logger: logr.Discard(), + RunnerConfig: runner.Config{ + AppName: appName, + Agent: agent, + }, + }) + reqCtx := &a2asrv.ExecutorContext{ + TaskID: "task-1", + ContextID: contextID, + Message: a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("hi")), + } + + var updates []*a2atype.TaskArtifactUpdateEvent + var completed *a2atype.TaskStatusUpdateEvent + for event, err := range executor.Execute(context.Background(), reqCtx) { + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + switch event := event.(type) { + case *a2atype.TaskArtifactUpdateEvent: + updates = append(updates, event) + case *a2atype.TaskStatusUpdateEvent: + if event.Status.State == a2atype.TaskStateCompleted { + completed = event + } + } + } + + if len(updates) != 2 { + t.Fatalf("artifact updates = %d, want 2", len(updates)) } - if ev.Status.Message.TaskID != a2atype.TaskID("task-xyz") { - t.Errorf("message TaskID = %q, want %q", ev.Status.Message.TaskID, a2atype.TaskID("task-xyz")) + if updates[0].Append || updates[0].LastChunk || updates[0].Artifact.Parts[0].Text() != "hel" { + t.Fatalf("first artifact update = %#v, want opening partial artifact", updates[0]) } - // The event itself also carries the ids (from reqCtx), matching the message. - if ev.ContextID != "ctx-xyz" { - t.Errorf("event ContextID = %q, want %q", ev.ContextID, "ctx-xyz") + if updates[1].Append || !updates[1].LastChunk || updates[1].Artifact.ID != updates[0].Artifact.ID || updates[1].Artifact.Parts[0].Text() != "hello" { + t.Fatalf("second artifact update = %#v, want content-bearing final replacement", updates[1]) } - if ev.TaskID != a2atype.TaskID("task-xyz") { - t.Errorf("event TaskID = %q, want %q", ev.TaskID, a2atype.TaskID("task-xyz")) + if completed == nil || completed.Status.Message != nil { + t.Fatalf("completed status = %#v, want content-free completion", completed) } } diff --git a/go/adk/pkg/models/openai_adk.go b/go/adk/pkg/models/openai_adk.go index f1bcce565d..89665545f1 100644 --- a/go/adk/pkg/models/openai_adk.go +++ b/go/adk/pkg/models/openai_adk.go @@ -378,13 +378,14 @@ func runStreaming(ctx context.Context, m *OpenAIModel, params openai.ChatComplet var aggregatedText strings.Builder toolCallsAcc := make(map[int64]map[string]any) var finishReason string - var promptTokens, completionTokens int64 + var promptTokens, completionTokens, totalTokens int64 for stream.Next() { chunk := stream.Current() if chunk.Usage.PromptTokens > 0 || chunk.Usage.CompletionTokens > 0 { promptTokens = chunk.Usage.PromptTokens completionTokens = chunk.Usage.CompletionTokens + totalTokens = chunk.Usage.TotalTokens } if len(chunk.Choices) == 0 { continue @@ -466,6 +467,7 @@ func runStreaming(ctx context.Context, m *OpenAIModel, params openai.ChatComplet usage = &genai.GenerateContentResponseUsageMetadata{ PromptTokenCount: int32(promptTokens), CandidatesTokenCount: int32(completionTokens), + TotalTokenCount: int32(totalTokens), } } resp := &model.LLMResponse{ @@ -527,6 +529,7 @@ func chatCompletionToLLMResponse(completion *openai.ChatCompletion) *model.LLMRe usage = &genai.GenerateContentResponseUsageMetadata{ PromptTokenCount: int32(completion.Usage.PromptTokens), CandidatesTokenCount: int32(completion.Usage.CompletionTokens), + TotalTokenCount: int32(completion.Usage.TotalTokens), } } return &model.LLMResponse{ diff --git a/go/adk/pkg/taskstore/store.go b/go/adk/pkg/taskstore/store.go index 4c52a0d21b..ed4709dec1 100644 --- a/go/adk/pkg/taskstore/store.go +++ b/go/adk/pkg/taskstore/store.go @@ -13,13 +13,9 @@ import ( a2ataskstore "github.com/a2aproject/a2a-go/v2/a2asrv/taskstore" ) -// Constants for partial-event metadata keys (inlined to avoid import cycle). const ( - metadataKeyKagentPartial = "kagent_partial" - metadataKeyKagentAdkPartial = "kagent_adk_partial" - metadataKeyAdkPartial = "adk_partial" - headerContentType = "Content-Type" - contentTypeJSON = "application/json" + headerContentType = "Content-Type" + contentTypeJSON = "application/json" ) // KAgentTaskStore persists A2A tasks to KAgent via REST API and implements @@ -48,64 +44,12 @@ type KAgentTaskResponse struct { Message string `json:"message,omitempty"` } -// isPartialMeta checks if a metadata map has a partial flag set to true. -// It checks the canonical kagent key (kagent_adk_partial) as well as legacy keys -// (adk_partial, kagent_partial) so that events from any prefix are recognised. -func isPartialMeta(meta map[string]any) bool { - if meta == nil { - return false - } - for _, key := range []string{metadataKeyKagentPartial, metadataKeyAdkPartial, metadataKeyKagentAdkPartial} { - if partial, ok := meta[key].(bool); ok && partial { - return true - } - } - return false -} - -// cleanPartialEvents removes partial streaming events from history. -func cleanPartialEvents(history []*a2atype.Message) []*a2atype.Message { - var cleaned []*a2atype.Message - for _, item := range history { - if item != nil && isPartialMeta(item.Metadata) { - continue - } - if item != nil && len(item.Parts) > 0 { - cleaned = append(cleaned, item) - } - } - return cleaned -} - -// cleanPartialArtifacts removes partial streaming artifacts. -func cleanPartialArtifacts(artifacts []*a2atype.Artifact) []*a2atype.Artifact { - var cleaned []*a2atype.Artifact - for _, a := range artifacts { - if a != nil && isPartialMeta(a.Metadata) { - continue - } - if a != nil && len(a.Parts) > 0 { - cleaned = append(cleaned, a) - } - } - return cleaned -} - func (s *KAgentTaskStore) saveTask(ctx context.Context, task *a2atype.Task) (a2ataskstore.TaskVersion, error) { if task == nil { return a2ataskstore.TaskVersionMissing, fmt.Errorf("task cannot be nil") } - // Work on a shallow copy so the caller's task is not mutated. - taskCopy := *task - if taskCopy.History != nil { - taskCopy.History = cleanPartialEvents(taskCopy.History) - } - if taskCopy.Artifacts != nil { - taskCopy.Artifacts = cleanPartialArtifacts(taskCopy.Artifacts) - } - - taskJSON, err := json.Marshal(&taskCopy) + taskJSON, err := json.Marshal(task) if err != nil { return a2ataskstore.TaskVersionMissing, fmt.Errorf("failed to marshal task: %w", err) } diff --git a/go/core/cli/internal/tui/chat.go b/go/core/cli/internal/tui/chat.go index 15e62cdfdf..a36c8c4e3d 100644 --- a/go/core/cli/internal/tui/chat.go +++ b/go/core/cli/internal/tui/chat.go @@ -44,6 +44,10 @@ type toolResult struct { Response any `json:"response"` } +type artifactBuffer struct { + text string +} + type chatModel struct { agentRef string sessionID string @@ -64,6 +68,9 @@ type chatModel struct { cancel context.CancelFunc streaming bool + artifacts map[a2atype.ArtifactID]*artifactBuffer + artifactOrder []a2atype.ArtifactID + showInput bool } @@ -94,6 +101,7 @@ func newChatModel(agentRef string, sessionID string, send SendMessageFn, verbose send: send, history: initial, spin: sp, + artifacts: make(map[a2atype.ArtifactID]*artifactBuffer), showInput: true, } } @@ -170,6 +178,7 @@ func (m *chatModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } case clia2a.StreamResult: if msg.Err != nil { + m.flushPendingArtifacts() m.appendError(msg.Err) m.streaming = false m.working = false @@ -179,6 +188,7 @@ func (m *chatModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.appendEvent(msg.Event) return m, m.waitNext() case streamDoneMsg: + m.flushPendingArtifacts() m.streaming = false m.working = false m.updateStatus() @@ -258,6 +268,7 @@ func (m *chatModel) appendEvent(ev a2atype.Event) { case *a2atype.TaskStatusUpdateEvent: final := res.Status.State.Terminal() if final { + m.flushPendingArtifacts() m.working = false m.updateStatus() } else if res.Status.Timestamp != nil { @@ -265,24 +276,27 @@ func (m *chatModel) appendEvent(ev a2atype.Event) { } else { m.setWorkingTime(time.Time{}) } - if res.Status.Message != nil { - m.handleMessageParts(res.Status.Message, final) - } + m.handleStatusMessage(res.Status.State, res.Status.Message) case *a2atype.TaskArtifactUpdateEvent: - if res.LastChunk { - text := extractTextFromParts(res.Artifact.Parts) - if strings.TrimSpace(text) != "" { - m.appendLine(theme.AgentStyle().Render("Agent:") + "\n" + text) - } - } + m.handleArtifactUpdate(res) case *a2atype.Message: m.handleMessageParts(res, true) case *a2atype.Task: - if len(res.History) > 0 { - last := res.History[len(res.History)-1] - m.handleMessageParts(last, true) - } + // A Task snapshot carries assembled output in artifacts. History is a + // collection of protocol Messages and must not be treated as task result. + for _, artifact := range res.Artifacts { + if artifact == nil { + continue + } + m.handleArtifactUpdate(&a2atype.TaskArtifactUpdateEvent{ + TaskID: res.ID, + ContextID: res.ContextID, + Artifact: artifact, + LastChunk: true, + }) + } + m.handleStatusMessage(res.Status.State, res.Status.Message) default: if m.verbose { if b, err := json.Marshal(ev); err == nil { @@ -292,6 +306,76 @@ func (m *chatModel) appendEvent(ev a2atype.Event) { } } +// handleStatusMessage processes control-plane status content only. Normal task +// output is delivered exclusively through artifacts. +func (m *chatModel) handleStatusMessage(state a2atype.TaskState, msg *a2atype.Message) { + if msg == nil { + return + } + switch state { + case a2atype.TaskStateInputRequired: + // Show tool/confirmation details but do not treat status text as output. + m.handleMessageParts(msg, false) + case a2atype.TaskStateAuthRequired, a2atype.TaskStateFailed: + m.handleMessageParts(msg, true) + } +} + +// handleArtifactUpdate merges text according to the A2A artifact update +// contract. Data parts are handled immediately so tool activity is visible +// even when an artifact has not reached its last chunk. +func (m *chatModel) handleArtifactUpdate(update *a2atype.TaskArtifactUpdateEvent) { + if update == nil || update.Artifact == nil { + return + } + + // handleMessageParts always processes tool parts; false suppresses text + // because text is committed only after the artifact has been assembled. + msg := a2atype.NewMessage(a2atype.MessageRoleAgent, update.Artifact.Parts...) + m.handleMessageParts(msg, false) + + text := extractTextFromParts(update.Artifact.Parts) + id := update.Artifact.ID + buffer, exists := m.artifacts[id] + if !exists { + buffer = &artifactBuffer{} + m.artifacts[id] = buffer + m.artifactOrder = append(m.artifactOrder, id) + } + if update.Append { + buffer.text += text + } else { + buffer.text = text + } + + if update.LastChunk { + m.commitArtifact(id) + } +} + +func (m *chatModel) commitArtifact(id a2atype.ArtifactID) { + buffer, ok := m.artifacts[id] + if !ok { + return + } + delete(m.artifacts, id) + for i, pendingID := range m.artifactOrder { + if pendingID == id { + m.artifactOrder = append(m.artifactOrder[:i], m.artifactOrder[i+1:]...) + break + } + } + if strings.TrimSpace(buffer.text) != "" { + m.appendLine(theme.AgentStyle().Render("Agent:") + "\n" + buffer.text) + } +} + +func (m *chatModel) flushPendingArtifacts() { + for len(m.artifactOrder) > 0 { + m.commitArtifact(m.artifactOrder[0]) + } +} + func (m *chatModel) appendError(err error) { m.appendLine(theme.ErrorStyle().Render(fmt.Sprintf("Error: %v", err))) } @@ -447,6 +531,8 @@ func (m *chatModel) appendLine(s string) { // ResetTranscript clears the viewport with a new header/title. func (m *chatModel) ResetTranscript(title string) { m.history = title + m.artifacts = make(map[a2atype.ArtifactID]*artifactBuffer) + m.artifactOrder = nil m.vp.SetContent(m.history) m.vp.GotoBottom() } @@ -464,12 +550,6 @@ func extractTextFromParts(parts a2atype.ContentParts) string { } if text := p.Text(); text != "" { b.WriteString(text) - continue - } - if data := p.Data(); data != nil { - if jp, err := json.Marshal(data); err == nil { - b.WriteString(string(jp)) - } } } return b.String() diff --git a/go/core/cli/internal/tui/chat_test.go b/go/core/cli/internal/tui/chat_test.go index cd8f766a65..e8958baeae 100644 --- a/go/core/cli/internal/tui/chat_test.go +++ b/go/core/cli/internal/tui/chat_test.go @@ -7,6 +7,7 @@ import ( "testing" a2atype "github.com/a2aproject/a2a-go/v2/a2a" + "github.com/a2aproject/a2a-go/v2/a2asrv" clia2a "github.com/kagent-dev/kagent/go/core/cli/internal/a2a" "github.com/stretchr/testify/require" ) @@ -29,3 +30,104 @@ func TestChatModelDisplaysStreamError(t *testing.T) { require.False(t, got.working) require.True(t, strings.Contains(got.history, "Error: stream disconnected")) } + +func TestChatModelBuffersArtifactDeltasUntilTerminalStatus(t *testing.T) { + model := newTestChatModel() + reqCtx := &a2asrv.ExecutorContext{TaskID: "task-1", ContextID: "ctx-1"} + first := a2atype.NewArtifactEvent(reqCtx, a2atype.NewTextPart("hel")) + second := a2atype.NewArtifactUpdateEvent(reqCtx, first.Artifact.ID, a2atype.NewTextPart("lo")) + + model.appendEvent(first) + model.appendEvent(second) + require.NotContains(t, model.history, "hello", "open artifacts should not be displayed before completion") + + model.appendEvent(a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateCompleted, nil)) + require.Contains(t, model.history, "hello") + require.Empty(t, model.artifacts) + require.Empty(t, model.artifactOrder) +} + +func TestChatModelArtifactReplacementAndContentBearingLastChunk(t *testing.T) { + model := newTestChatModel() + reqCtx := &a2asrv.ExecutorContext{TaskID: "task-1", ContextID: "ctx-1"} + partial := a2atype.NewArtifactEvent(reqCtx, a2atype.NewTextPart("hel")) + final := a2atype.NewArtifactUpdateEvent(reqCtx, partial.Artifact.ID, a2atype.NewTextPart("hello")) + final.Append = false + final.LastChunk = true + + model.appendEvent(partial) + model.appendEvent(final) + + require.Contains(t, model.history, "hello") + require.NotContains(t, model.history, "helhello", "append=false must replace buffered partial text") + require.Empty(t, model.artifacts) + + // A later terminal status must not display an already-closed artifact again. + before := model.history + model.appendEvent(a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateCompleted, nil)) + require.Equal(t, before, model.history) +} + +func TestChatModelProcessesToolPartsBeforeLastChunk(t *testing.T) { + model := newTestChatModel() + reqCtx := &a2asrv.ExecutorContext{TaskID: "task-1", ContextID: "ctx-1"} + callPart := a2atype.NewDataPart(map[string]any{ + "name": "get_pods", + "id": "call-1", + "args": map[string]any{"namespace": "default"}, + }) + callPart.Metadata = map[string]any{"adk_type": "function_call"} + callUpdate := a2atype.NewArtifactEvent(reqCtx, callPart) + resultPart := a2atype.NewDataPart(map[string]any{ + "name": "get_pods", + "id": "call-1", + "response": map[string]any{"pods": []any{"pod-a"}}, + }) + resultPart.Metadata = map[string]any{"adk_type": "function_response"} + resultUpdate := a2atype.NewArtifactEvent(reqCtx, resultPart) + + model.appendEvent(callUpdate) + model.appendEvent(resultUpdate) + + require.Contains(t, model.history, "Tool Call: get_pods") + require.Contains(t, model.history, "Tool Result: get_pods") + require.Contains(t, model.history, "call-1") + require.Contains(t, model.history, "pod-a") +} + +func TestChatModelPreservesFailedStatusExplanation(t *testing.T) { + model := newTestChatModel() + reqCtx := &a2asrv.ExecutorContext{TaskID: "task-1", ContextID: "ctx-1"} + message := a2atype.NewMessage(a2atype.MessageRoleAgent, a2atype.NewTextPart("execution failed")) + + model.appendEvent(a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateFailed, message)) + + require.Contains(t, model.history, "execution failed") +} + +func TestChatModelReadsTaskSnapshotOutputFromArtifactsOnly(t *testing.T) { + model := newTestChatModel() + artifact := &a2atype.Artifact{ + ID: "artifact-1", + Parts: a2atype.ContentParts{a2atype.NewTextPart("artifact result")}, + } + task := &a2atype.Task{ + ID: "task-1", + ContextID: "ctx-1", + Status: a2atype.TaskStatus{State: a2atype.TaskStateCompleted}, + Artifacts: []*a2atype.Artifact{artifact}, + } + + model.appendEvent(task) + + require.Contains(t, model.history, "artifact result") +} + +func newTestChatModel() *chatModel { + send := func(context.Context, *a2atype.SendMessageRequest) <-chan clia2a.StreamResult { + ch := make(chan clia2a.StreamResult) + close(ch) + return ch + } + return newChatModel("default/agent", "session-1", send, false) +} diff --git a/go/core/test/e2e/foundry_test.go b/go/core/test/e2e/foundry_test.go index 72bf149e5c..9d572b98ba 100644 --- a/go/core/test/e2e/foundry_test.go +++ b/go/core/test/e2e/foundry_test.go @@ -153,7 +153,6 @@ func TestE2EMemoryWithGoADKFoundryAgent(t *testing.T) { saveResult = runSyncTest(t, a2aClient, "Remember that I prefer dark mode and Go over Python", "saved your preferences to memory", - nil, ) }) @@ -161,7 +160,6 @@ func TestE2EMemoryWithGoADKFoundryAgent(t *testing.T) { runSyncTest(t, a2aClient, "What are my preferences?", "dark mode", - nil, saveResult.ContextID, ) }) diff --git a/go/core/test/e2e/invoke_api_test.go b/go/core/test/e2e/invoke_api_test.go index 3e1fcf76f2..54a34f20a8 100644 --- a/go/core/test/e2e/invoke_api_test.go +++ b/go/core/test/e2e/invoke_api_test.go @@ -280,10 +280,9 @@ var defaultRetry = wait.Backoff{ Jitter: 0.2, } -// runSyncTest runs a synchronous message test -// useArtifacts: if true, check artifacts; if false or nil, check history; +// runSyncTest runs a synchronous message test and validates task artifact output. // contextID: optional context ID to maintain conversation context -func runSyncTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, expectedText string, useArtifacts *bool, contextID ...string) *a2atype.Task { +func runSyncTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, expectedText string, contextID ...string) *a2atype.Task { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() @@ -313,25 +312,17 @@ func runSyncTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, expecte taskResult, ok := result.(*a2atype.Task) require.True(t, ok) - // Extract text based on useArtifacts flag - if useArtifacts != nil && *useArtifacts { - // Check artifacts (used by CrewAI flows) - text := extractTextFromArtifacts(taskResult) - require.Contains(t, text, expectedText) - } else { - // Check history (used by declarative agents) - default - text := a2a.ExtractText(taskResult.History[len(taskResult.History)-1]) - jsn, err := json.Marshal(taskResult) - require.NoError(t, err) - require.Contains(t, text, expectedText, string(jsn)) - } + text := extractTextFromArtifacts(taskResult) + jsn, err := json.Marshal(taskResult) + require.NoError(t, err) + require.Contains(t, text, expectedText, string(jsn)) return taskResult } // runStreamingTest runs a streaming message test // If contextID is provided, it will be included in the message to maintain conversation context -// Checks the full JSON output to support both artifacts and history from different agent types +// The last completed artifact contains the final agent output. func runStreamingTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, expectedText string, contextID ...string) { msg := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart(userMessage)) @@ -352,7 +343,8 @@ func runStreamingTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, ex t.Logf("%s trying to open stream", time.Now().Format(time.RFC3339)) stream := a2aClient.SendStreamingMessage(ctx, &a2atype.SendMessageRequest{Message: msg}) - texts := make([]string, 0) + lastText = "" + foundFinalArtifact := false eventCount := 0 for event, streamErr := range stream { if streamErr != nil { @@ -363,9 +355,14 @@ func runStreamingTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, ex if event == nil { continue } - texts = append(texts, extractTextFromEvent(event)) + if artifactUpdate, ok := event.(*a2atype.TaskArtifactUpdateEvent); ok && artifactUpdate.LastChunk && artifactUpdate.Artifact != nil { + lastText = a2a.ExtractText(&a2atype.Message{Parts: artifactUpdate.Artifact.Parts}) + foundFinalArtifact = true + } + } + if !foundFinalArtifact { + return fmt.Errorf("streaming response contained no completed artifact (%d events)", eventCount) } - lastText = strings.Join(texts, "\n") if !strings.Contains(lastText, expectedText) { t.Logf("%s stream completed but expected text %q not found in response (got %d events)", time.Now().Format(time.RFC3339), expectedText, eventCount) @@ -378,26 +375,6 @@ func runStreamingTest(t *testing.T, a2aClient *a2aclient.Client, userMessage, ex require.NoError(t, err, lastText) } -func extractTextFromEvent(event a2atype.Event) string { - switch e := event.(type) { - case *a2atype.TaskStatusUpdateEvent: - return a2a.ExtractText(e.Status.Message) - case *a2atype.TaskArtifactUpdateEvent: - return a2a.ExtractText(&a2atype.Message{Parts: e.Artifact.Parts}) - case *a2atype.Message: - return a2a.ExtractText(e) - case *a2atype.Task: - text := strings.Builder{} - if e.Status.Message != nil { - text.WriteString(a2a.ExtractText(e.Status.Message)) - } - text.WriteString(extractTextFromArtifacts(e)) - return text.String() - default: - return "" - } -} - func a2aURL(namespace, name string, sandbox bool) string { kagentURL := os.Getenv("KAGENT_URL") if kagentURL == "" { @@ -609,7 +586,7 @@ func TestE2EInvokeInlineAgent(t *testing.T) { // Run tests t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "List all nodes in the cluster", "kagent-control-plane", nil) + runSyncTest(t, a2aClient, "List all nodes in the cluster", "kagent-control-plane") }) t.Run("streaming_invocation", func(t *testing.T) { @@ -706,7 +683,7 @@ func TestE2EInvokeExternalAgent(t *testing.T) { // Run tests t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "What can you do?", "kebab", nil) + runSyncTest(t, a2aClient, "What can you do?", "kebab") }) t.Run("streaming_invocation", func(t *testing.T) { @@ -717,7 +694,7 @@ func TestE2EInvokeExternalAgent(t *testing.T) { // Setup A2A client with authentication authClient := newA2AClient(t, a2aURL, nil, map[string]string{"x-user-id": "user@example.com"}) - runSyncTest(t, authClient, "What can you do?", "kebab for user@example.com", nil) + runSyncTest(t, authClient, "What can you do?", "kebab for user@example.com") }) } @@ -754,7 +731,7 @@ func TestE2EInvokeDeclarativeAgentWithMcpServerTool(t *testing.T) { // Run tests t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "add 3 and 5", "8", nil) + runSyncTest(t, a2aClient, "add 3 and 5", "8") }) t.Run("streaming_invocation", func(t *testing.T) { @@ -915,9 +892,8 @@ func TestE2EInvokeOpenAIAgent(t *testing.T) { a2aURL := a2aUrl("kagent", "basic-openai-test-agent") a2aClient := newA2AClient(t, a2aURL, nil, nil) - useArtifacts := true t.Run("sync_invocation_calculator", func(t *testing.T) { - runSyncTest(t, a2aClient, "What is 2+2?", "4", &useArtifacts) + runSyncTest(t, a2aClient, "What is 2+2?", "4") }) t.Run("streaming_invocation_weather", func(t *testing.T) { @@ -978,7 +954,7 @@ func TestE2EInvokeLangGraphAgent(t *testing.T) { a2aClient := newA2AClient(t, a2aURL, nil, nil) t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "make me a kebab", "kebab is ready", nil) + runSyncTest(t, a2aClient, "make me a kebab", "kebab is ready") }) t.Run("streaming_invocation", func(t *testing.T) { @@ -1052,13 +1028,11 @@ func TestE2EInvokeCrewAIAgent(t *testing.T) { t.Run("two_turn_conversation", func(t *testing.T) { // First turn: Generate initial poem - // Use artifacts only (true) for CrewAI flows - useArtifacts := true - taskResult1 := runSyncTest(t, a2aClient, "Generate a poem about CrewAI", "CrewAI is awesome, it makes coding fun.", &useArtifacts) + taskResult1 := runSyncTest(t, a2aClient, "Generate a poem about CrewAI", "CrewAI is awesome, it makes coding fun.") // Second turn: Continue poem (tests persistence) // Use the same ContextID to maintain conversation context - runSyncTest(t, a2aClient, "Continue the poem", "In harmony with the code, it flows so smooth.", &useArtifacts, taskResult1.ContextID) + runSyncTest(t, a2aClient, "Continue the poem", "In harmony with the code, it flows so smooth.", taskResult1.ContextID) }) t.Run("streaming_invocation", func(t *testing.T) { @@ -1144,7 +1118,7 @@ func runE2EInvokeSTSIntegration(t *testing.T, runtimeName string, runtimeOverrid a2aClient := newA2AClient(t, a2aURL, httpClient, nil) t.Run(runtimeName+"/sts_exchange_sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "add 3 and 5", "8", nil) + runSyncTest(t, a2aClient, "add 3 and 5", "8") // verify our mock STS server received the token exchange request stsRequests := stsServer.GetRequests() @@ -1182,7 +1156,7 @@ func TestE2EInvokeSkillInAgent(t *testing.T) { a2aClient := setupA2AClient(t, agent) // Run tests - runSyncTest(t, a2aClient, "make me a kebab", "Pick it up from around the corner", nil) + runSyncTest(t, a2aClient, "make me a kebab", "Pick it up from around the corner") } func TestE2ESkillImagePullSecrets(t *testing.T) { @@ -1259,7 +1233,7 @@ func TestE2ESkillImagePullSecrets(t *testing.T) { // Verify the agent works end-to-end with the skill a2aClient := setupA2AClient(t, agent) - runSyncTest(t, a2aClient, "make me a kebab", "Pick it up from around the corner", nil) + runSyncTest(t, a2aClient, "make me a kebab", "Pick it up from around the corner") } func TestE2EDeclarativeAgentNetworkAllowlistWithSkills(t *testing.T) { @@ -1289,7 +1263,7 @@ func runDeclarativeAgentNetworkAllowlistWithSkills(t *testing.T, runtimeName str }) a2aClient := setupA2AClient(t, agent) - runSyncTest(t, a2aClient, "check the controller health with bash", "python and node are available; network denied", nil) + runSyncTest(t, a2aClient, "check the controller health with bash", "python and node are available; network denied") }) t.Run(runtimeName+"/allowlist_enables_access", func(t *testing.T) { @@ -1307,7 +1281,7 @@ func runDeclarativeAgentNetworkAllowlistWithSkills(t *testing.T, runtimeName str }) a2aClient := setupA2AClient(t, agent) - runSyncTest(t, a2aClient, "check the controller health with bash", "python and node are available; controller health is ok", nil) + runSyncTest(t, a2aClient, "check the controller health with bash", "python and node are available; controller health is ok") }) } @@ -1362,7 +1336,7 @@ func TestE2EInvokePassthroughAgent(t *testing.T) { // Authorization header "Bearer passthrough-test-token-12345". // If passthrough is broken, mockllm returns 404 and the test fails. t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "Hello from passthrough", "Token received successfully via passthrough", nil) + runSyncTest(t, a2aClient, "Hello from passthrough", "Token received successfully via passthrough") }) } @@ -1382,7 +1356,7 @@ func TestE2EAgentDefaultRuntimeIsGo(t *testing.T) { a2aClient := setupA2AClient(t, agent) t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "What is 2+2?", "4", nil) + runSyncTest(t, a2aClient, "What is 2+2?", "4") }) t.Run("streaming_invocation", func(t *testing.T) { @@ -1416,7 +1390,6 @@ func runMemoryAgentTest(t *testing.T, extraOpts AgentOptions) { saveResult = runSyncTest(t, a2aClient, "Remember that I prefer dark mode and Go over Python", "saved your preferences to memory", - nil, ) }) @@ -1424,7 +1397,6 @@ func runMemoryAgentTest(t *testing.T, extraOpts AgentOptions) { runSyncTest(t, a2aClient, "What are my preferences?", "dark mode", - nil, saveResult.ContextID, ) }) @@ -1509,7 +1481,7 @@ You are {{.AgentName}}, operating in {{.AgentNamespace}}. a2aClient := setupA2AClient(t, agent) t.Run("sync_invocation", func(t *testing.T) { - runSyncTest(t, a2aClient, "List all nodes in the cluster", "kagent-control-plane", nil) + runSyncTest(t, a2aClient, "List all nodes in the cluster", "kagent-control-plane") }) t.Run("streaming_invocation", func(t *testing.T) { @@ -1545,7 +1517,7 @@ You are {{.AgentName}}, operating in {{.AgentNamespace}}. } // Verify the agent still responds correctly - runSyncTest(t, a2aClient, "List all nodes in the cluster", "kagent-control-plane", nil) + runSyncTest(t, a2aClient, "List all nodes in the cluster", "kagent-control-plane") }) } @@ -1568,7 +1540,7 @@ func TestE2EIAgentRunsCode(t *testing.T) { a2aClient := setupA2AClient(t, agent) // Run tests - runSyncTest(t, a2aClient, "write some code", "hello, world!", nil) + runSyncTest(t, a2aClient, "write some code", "hello, world!") } func cleanup(t *testing.T, cli client.Client, obj ...client.Object) { diff --git a/go/core/test/e2e/remotemcpserver_tls_test.go b/go/core/test/e2e/remotemcpserver_tls_test.go index c05cc64b52..fbedd029fc 100644 --- a/go/core/test/e2e/remotemcpserver_tls_test.go +++ b/go/core/test/e2e/remotemcpserver_tls_test.go @@ -367,7 +367,7 @@ func TestE2E_RMS_PrivateCAUpstream(t *testing.T) { }}) a2aClient := setupA2AClient(t, agent) - runSyncTest(t, a2aClient, "add 2 and 3", "5", nil) + runSyncTest(t, a2aClient, "add 2 and 3", "5") // The agent's tools/call should also have reached mockmcp. postInvoke := mcp.server.Requests() @@ -416,7 +416,7 @@ func TestE2E_RMS_DisableVerify(t *testing.T) { }}) a2aClient := setupA2AClient(t, agent) - runSyncTest(t, a2aClient, "add 2 and 3", "5", nil) + runSyncTest(t, a2aClient, "add 2 and 3", "5") } // TestE2E_RMS_SSE_TLS exercises the SSE-transport-with-TLS code path. @@ -472,7 +472,7 @@ func TestE2E_RMS_SSE_TLS(t *testing.T) { }}) a2aClient := setupA2AClient(t, agent) - runSyncTest(t, a2aClient, "add 2 and 3", "5", nil) + runSyncTest(t, a2aClient, "add 2 and 3", "5") } // TestE2E_API_ToolServerCompanionSecrets posts a ToolServerCreateRequest diff --git a/python/packages/kagent-adk/src/kagent/adk/_agent_executor.py b/python/packages/kagent-adk/src/kagent/adk/_agent_executor.py index 85b1e0ee8e..fcc1f953e2 100644 --- a/python/packages/kagent-adk/src/kagent/adk/_agent_executor.py +++ b/python/packages/kagent-adk/src/kagent/adk/_agent_executor.py @@ -11,7 +11,6 @@ from a2a.server.agent_execution.context import RequestContext from a2a.server.events.event_queue_v2 import EventQueue from a2a.types import ( - Artifact, Message, Part, Role, @@ -22,16 +21,16 @@ TaskStatusUpdateEvent, ) from google.adk.events import Event, EventActions -from google.adk.flows.llm_flows.functions import REQUEST_CONFIRMATION_FUNCTION_CALL_NAME +from google.adk.flows.llm_flows.functions import REQUEST_CONFIRMATION_FUNCTION_CALL_NAME, REQUEST_EUC_FUNCTION_CALL_NAME from google.adk.runners import Runner from google.adk.sessions import Session from google.adk.tools.tool_confirmation import ToolConfirmation from google.adk.utils.context_utils import Aclosing from google.genai import types as genai_types +from google.protobuf.json_format import MessageToDict from kagent.core.a2a import ( KAGENT_HITL_DECISION_TYPE_APPROVE, KAGENT_HITL_DECISION_TYPE_BATCH, - TaskResultAggregator, extract_ask_user_answers_from_message, extract_batch_decisions_from_message, extract_decision_from_message, @@ -50,6 +49,37 @@ logger = logging.getLogger("kagent_adk." + __name__) +def _is_long_running_function_call(part: Part) -> bool: + """True when a DataPart is a long-running function_call (HITL/auth).""" + if not part.HasField("data"): + return False + metadata = MessageToDict(part.metadata) if part.metadata else {} + return bool(metadata.get(get_kagent_metadata_key("is_long_running"))) + + +def _split_hitl_artifact_parts( + event: TaskArtifactUpdateEvent, + hitl_parts: list[Part], +) -> TaskArtifactUpdateEvent | None: + """Move long-running function_call parts onto the HITL status; keep the rest. + + Mirrors Go adka2a inputRequiredProcessor: confirmation/auth parts belong on + input-required/auth-required status, while ordinary text/tool output stays + on the artifact stream. + """ + output_parts: list[Part] = [] + for part in event.artifact.parts: + if _is_long_running_function_call(part): + hitl_parts.append(part) + else: + output_parts.append(part) + if not output_parts: + return None + del event.artifact.parts[:] + event.artifact.parts.extend(output_parts) + return event + + class A2aAgentExecutorConfig(BaseModel): """Configuration for the KAgent A2aAgentExecutor.""" @@ -63,7 +93,7 @@ class A2aAgentExecutor(AgentExecutor): - Per-request runner lifecycle (created fresh and closed after each request) - OpenTelemetry span attribute management - Enhanced error handling (Ollama-specific JSON parse errors, CancelledError) - - Partial event filtering to avoid duplicate aggregation during streaming + - A2A artifact streaming with kagent HITL status handling - Session naming from first message text - Request header forwarding to session state - Invocation ID tracking in final event metadata @@ -544,7 +574,9 @@ async def _handle_request( if isinstance(tool, SubagentSessionProvider) and tool.subagent_session_id: subagent_session_ids[tool.name] = tool.subagent_session_id - task_result_aggregator = TaskResultAggregator() + hitl_parts: list[Part] = [] + terminal_status: TaskStatusUpdateEvent | None = None + agents_artifacts: dict[str, str] = {} async with Aclosing(runner.run_async(**run_args)) as agen: async for adk_event in agen: # Capture the real invocation_id from the first ADK event that has one @@ -559,21 +591,34 @@ async def _handle_request( if getattr(adk_event, "usage_metadata", None) is not None: last_usage_metadata = adk_event.usage_metadata - for a2a_event in convert_event_to_a2a_events( + a2a_events = convert_event_to_a2a_events( adk_event, invocation_context, context.task_id, context.context_id, subagent_session_ids=subagent_session_ids or None, - ): - # Only aggregate non-partial events to avoid duplicates from streaming chunks - # Partial events are sent to frontend for display but not accumulated - if not adk_event.partial: - task_result_aggregator.process_event(a2a_event) + agents_artifacts=agents_artifacts, + ) + + is_long_running = bool(getattr(adk_event, "long_running_tool_ids", None)) + for a2a_event in a2a_events: + if isinstance(a2a_event, TaskStatusUpdateEvent) and a2a_event.status.state in ( + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_AUTH_REQUIRED, + TaskState.TASK_STATE_INPUT_REQUIRED, + ): + terminal_status = a2a_event + elif is_long_running and isinstance(a2a_event, TaskArtifactUpdateEvent): + a2a_event = _split_hitl_artifact_parts(a2a_event, hitl_parts) + if a2a_event is None: + continue await event_queue.enqueue_event(a2a_event) + if terminal_status is not None: + break + # Break on confirmation events that use long running tools - if getattr(adk_event, "long_running_tool_ids", None): + if is_long_running: break # Attach the last LLM usage to run_metadata so the A2A task_manager @@ -581,45 +626,40 @@ async def _handle_request( if last_usage_metadata is not None: run_metadata[get_kagent_metadata_key("usage_metadata")] = serialize_metadata_value(last_usage_metadata) - # publish the task result event - this is final - if ( - task_result_aggregator.task_state == TaskState.TASK_STATE_WORKING - and task_result_aggregator.task_status_message is not None - and task_result_aggregator.task_status_message.parts - ): - # if task is still working properly, publish the artifact update event as - # the final result according to a2a protocol. - await event_queue.enqueue_event( - TaskArtifactUpdateEvent( - task_id=context.task_id, - last_chunk=True, - context_id=context.context_id, - artifact=Artifact( - artifact_id=str(uuid.uuid4()), - parts=task_result_aggregator.task_status_message.parts, - ), - ) - ) - # publish the final status update event + if hitl_parts: + hitl_state = TaskState.TASK_STATE_INPUT_REQUIRED + for part in hitl_parts: + if not part.HasField("data"): + continue + payload = MessageToDict(part.data) + if isinstance(payload, dict) and payload.get("name") == REQUEST_EUC_FUNCTION_CALL_NAME: + hitl_state = TaskState.TASK_STATE_AUTH_REQUIRED + break await event_queue.enqueue_event( TaskStatusUpdateEvent( task_id=context.task_id, + context_id=context.context_id, status=TaskStatus( - state=TaskState.TASK_STATE_COMPLETED, + state=hitl_state, timestamp=now_timestamp(), + message=Message( + message_id=str(uuid.uuid4()), + role=Role.ROLE_AGENT, + parts=hitl_parts, + task_id=context.task_id, + context_id=context.context_id, + ), ), - context_id=context.context_id, metadata=run_metadata, ) ) - else: + elif terminal_status is None: await event_queue.enqueue_event( TaskStatusUpdateEvent( task_id=context.task_id, status=TaskStatus( - state=task_result_aggregator.task_state, + state=TaskState.TASK_STATE_COMPLETED, timestamp=now_timestamp(), - message=task_result_aggregator.task_status_message, ), context_id=context.context_id, metadata=run_metadata, diff --git a/python/packages/kagent-adk/src/kagent/adk/converters/event_converter.py b/python/packages/kagent-adk/src/kagent/adk/converters/event_converter.py index 10a8d68fbc..d59d07344c 100644 --- a/python/packages/kagent-adk/src/kagent/adk/converters/event_converter.py +++ b/python/packages/kagent-adk/src/kagent/adk/converters/event_converter.py @@ -5,11 +5,19 @@ from typing import Any, Dict, List, Optional from a2a.server.events import Event as A2AEvent -from a2a.types import Message, Role, Task, TaskState, TaskStatus, TaskStatusUpdateEvent +from a2a.types import ( + Artifact, + Message, + Role, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) from a2a.types import Part as A2APart from google.adk.agents.invocation_context import InvocationContext from google.adk.events.event import Event -from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME from google.genai import types as genai_types from google.protobuf.json_format import MessageToDict from kagent.core.a2a import ( @@ -78,11 +86,7 @@ def _get_context_metadata(event: Event, invocation_context: InvocationContext) - raise ValueError("Invocation context cannot be None") try: - partial = getattr(event, "partial", False) - if not isinstance(partial, bool): - partial = False metadata: Dict[str, Any] = { - get_kagent_metadata_key("adk_partial"): partial, get_kagent_metadata_key("app_name"): invocation_context.app_name, get_kagent_metadata_key("user_id"): invocation_context.user_id, get_kagent_metadata_key("session_id"): invocation_context.session.id, @@ -286,14 +290,15 @@ def _create_error_status_event( ) -def _create_status_update_event( +def _create_artifact_update_event( message: Message, invocation_context: InvocationContext, event: Event, task_id: Optional[str] = None, context_id: Optional[str] = None, -) -> TaskStatusUpdateEvent: - """Creates a TaskStatusUpdateEvent for running scenarios. + agents_artifacts: Optional[Dict[str, str]] = None, +) -> Optional[TaskArtifactUpdateEvent]: + """Creates a TaskArtifactUpdateEvent for task output. Args: message: The A2A message to include. @@ -304,43 +309,37 @@ def _create_status_update_event( Returns: - A TaskStatusUpdateEvent with RUNNING state. + A TaskArtifactUpdateEvent containing the converted output parts. """ - status = TaskStatus( - state=TaskState.TASK_STATE_WORKING, - message=message, - timestamp=now_timestamp(), - ) - - has_auth_required = False - has_input_required = False - for part in message.parts: - if not part.HasField("data"): - continue - metadata = MessageToDict(part.metadata) if part.metadata else {} - if ( - metadata.get(get_kagent_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY)) - != A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL - ): - continue - if metadata.get(get_kagent_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)) is not True: - continue - payload = MessageToDict(part.data) - if isinstance(payload, dict) and payload.get("name") == REQUEST_EUC_FUNCTION_CALL_NAME: - has_auth_required = True - break - has_input_required = True - - if has_auth_required: - status.state = TaskState.TASK_STATE_AUTH_REQUIRED - elif has_input_required: - status.state = TaskState.TASK_STATE_INPUT_REQUIRED - - return TaskStatusUpdateEvent( + metadata = _get_context_metadata(event, invocation_context) + partial = bool(getattr(event, "partial", False)) + # Match Go adka2a.OutputArtifactPerEvent: reuse one artifact ID across + # partial deltas, append while partial, then replace+close on the final + # non-partial event (which may repeat the full text). + artifact_id = str(uuid.uuid4()) + append = False + if agents_artifacts is not None: + agent_name = event.author or "" + active_artifact_id = agents_artifacts.get(agent_name) + if active_artifact_id: + artifact_id = active_artifact_id + append = partial + if partial: + agents_artifacts[agent_name] = artifact_id + elif active_artifact_id: + del agents_artifacts[agent_name] + + return TaskArtifactUpdateEvent( task_id=task_id, context_id=context_id, - status=status, - metadata=_get_context_metadata(event, invocation_context), + append=append, + last_chunk=not partial, + artifact=Artifact( + artifact_id=artifact_id, + parts=list(message.parts), + metadata=metadata, + ), + metadata=metadata, ) @@ -350,6 +349,7 @@ def convert_event_to_a2a_events( task_id: Optional[str] = None, context_id: Optional[str] = None, subagent_session_ids: Optional[Dict[str, str]] = None, + agents_artifacts: Optional[Dict[str, str]] = None, ) -> List[A2AEvent]: """Converts a GenAI event to a list of A2A events. @@ -360,6 +360,8 @@ def convert_event_to_a2a_events( context_id: Optional Context ID to use for generated events. subagent_session_ids: Optional mapping of tool name to pre-generated subagent session ID, threaded to ``convert_event_to_a2a_message``. + agents_artifacts: Mutable mapping used to reuse artifact IDs across + partial chunks from the same agent. Returns: A list of A2A events representing the converted ADK event. @@ -379,6 +381,7 @@ def convert_event_to_a2a_events( if event.error_code and not _is_normal_completion(event.error_code): error_event = _create_error_status_event(event, invocation_context, task_id, context_id) a2a_events.append(error_event) + return a2a_events # Handle regular message content message = convert_event_to_a2a_message( @@ -389,8 +392,16 @@ def convert_event_to_a2a_events( context_id=context_id, ) if message: - running_event = _create_status_update_event(message, invocation_context, event, task_id, context_id) - a2a_events.append(running_event) + artifact_event = _create_artifact_update_event( + message, + invocation_context, + event, + task_id, + context_id, + agents_artifacts, + ) + if artifact_event is not None: + a2a_events.append(artifact_event) except Exception as e: logger.error("Failed to convert event to A2A events: %s", e) diff --git a/python/packages/kagent-adk/tests/unittests/converters/test_event_converter.py b/python/packages/kagent-adk/tests/unittests/converters/test_event_converter.py index 753e1ffb96..115eb6c210 100644 --- a/python/packages/kagent-adk/tests/unittests/converters/test_event_converter.py +++ b/python/packages/kagent-adk/tests/unittests/converters/test_event_converter.py @@ -2,8 +2,9 @@ from unittest.mock import Mock import pytest -from a2a.types import TaskState, TaskStatusUpdateEvent +from a2a.types import TaskArtifactUpdateEvent, TaskState, TaskStatusUpdateEvent from google.genai import types as genai_types +from google.protobuf.json_format import MessageToDict from kagent.core.a2a import get_kagent_metadata_key from pydantic import BaseModel, Field @@ -19,7 +20,9 @@ def _create_mock_invocation_context(): return context -def _create_mock_event(error_code=None, content=None, invocation_id="test_invocation", author="test_author"): +def _create_mock_event( + error_code=None, content=None, invocation_id="test_invocation", author="test_author", partial=False +): """Create a mock event for testing.""" event = Mock() event.error_code = error_code @@ -31,6 +34,8 @@ def _create_mock_event(error_code=None, content=None, invocation_id="test_invoca event.custom_metadata = None event.usage_metadata = None event.error_message = None + event.partial = partial + event.long_running_tool_ids = None return event @@ -126,23 +131,91 @@ def test_convert_event_to_a2a_events(self): assert error_code_key in error_event.metadata assert error_event.metadata[error_code_key] == str(genai_types.FinishReason.MALFORMED_FUNCTION_CALL) - def test_message_carries_task_and_context_ids(self): - """The converted message stamps task_id/context_id so consumers that - flatten task.history can key it to its task without backfilling.""" + def test_content_is_emitted_as_artifact(self): invocation_context = _create_mock_invocation_context() content = genai_types.Content(parts=[genai_types.Part(text="hello world")]) event = _create_mock_event(content=content, invocation_id="test_invocation_ids") result = convert_event_to_a2a_events(event, invocation_context, task_id="task-xyz", context_id="ctx-xyz") - working_events = [ - e for e in result if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_WORKING - ] - assert len(working_events) == 1 - message = working_events[0].status.message - assert message is not None - assert message.task_id == "task-xyz" - assert message.context_id == "ctx-xyz" + artifact_events = [e for e in result if isinstance(e, TaskArtifactUpdateEvent)] + assert len(artifact_events) == 1 + artifact_event = artifact_events[0] + assert artifact_event.task_id == "task-xyz" + assert artifact_event.context_id == "ctx-xyz" + assert artifact_event.artifact.parts[0].text == "hello world" + assert artifact_event.last_chunk is True + assert get_kagent_metadata_key("adk_partial") not in artifact_event.metadata + assert not any( + isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_WORKING for e in result + ) + + def test_partial_chunks_reuse_artifact_id_and_final_replaces(self): + """Go OutputArtifactPerEvent framing: append deltas, replace+close on final.""" + invocation_context = _create_mock_invocation_context() + agents_artifacts: dict[str, str] = {} + + first = convert_event_to_a2a_events( + _create_mock_event(content=genai_types.Content(parts=[genai_types.Part(text="hel")]), partial=True), + invocation_context, + agents_artifacts=agents_artifacts, + )[0] + second = convert_event_to_a2a_events( + _create_mock_event(content=genai_types.Content(parts=[genai_types.Part(text="lo")]), partial=True), + invocation_context, + agents_artifacts=agents_artifacts, + )[0] + final = convert_event_to_a2a_events( + _create_mock_event(content=genai_types.Content(parts=[genai_types.Part(text="hello")]), partial=False), + invocation_context, + agents_artifacts=agents_artifacts, + )[0] + + assert isinstance(first, TaskArtifactUpdateEvent) + assert first.artifact.artifact_id == second.artifact.artifact_id == final.artifact.artifact_id + assert first.append is False + assert first.last_chunk is False + assert second.append is True + assert second.last_chunk is False + assert final.append is False + assert final.last_chunk is True + assert final.artifact.parts[0].text == "hello" + assert agents_artifacts == {} + + def test_final_mixed_event_keeps_text_and_hitl_parts_on_same_artifact(self): + invocation_context = _create_mock_invocation_context() + agents_artifacts: dict[str, str] = {} + partial_event = _create_mock_event( + content=genai_types.Content(parts=[genai_types.Part(text="partial text")]), partial=True + ) + partial_artifact = convert_event_to_a2a_events( + partial_event, invocation_context, agents_artifacts=agents_artifacts + )[0] + + final_event = _create_mock_event( + content=genai_types.Content( + parts=[ + genai_types.Part(text="partial text complete"), + genai_types.Part( + function_call=genai_types.FunctionCall(id="call-1", name="dangerous_tool", args={"value": "x"}) + ), + ] + ), + partial=False, + ) + final_event.long_running_tool_ids = {"call-1"} + + result = convert_event_to_a2a_events(final_event, invocation_context, agents_artifacts=agents_artifacts) + + assert len(result) == 1 + final_artifact = result[0] + assert isinstance(final_artifact, TaskArtifactUpdateEvent) + assert final_artifact.last_chunk is True + assert final_artifact.append is False + assert final_artifact.artifact.artifact_id == partial_artifact.artifact.artifact_id + assert final_artifact.artifact.parts[0].text == "partial text complete" + assert MessageToDict(final_artifact.artifact.parts[1].data)["name"] == "dangerous_tool" + assert agents_artifacts == {} class TestSerializeMetadataValue: diff --git a/python/packages/kagent-adk/tests/unittests/test_artifact_streaming.py b/python/packages/kagent-adk/tests/unittests/test_artifact_streaming.py new file mode 100644 index 0000000000..13ce2643fa --- /dev/null +++ b/python/packages/kagent-adk/tests/unittests/test_artifact_streaming.py @@ -0,0 +1,49 @@ +from a2a.types import Artifact, Part, TaskArtifactUpdateEvent +from google.protobuf.json_format import ParseDict +from google.protobuf.struct_pb2 import Value + +from kagent.adk._agent_executor import _split_hitl_artifact_parts +from kagent.core.a2a import get_kagent_metadata_key + + +def test_split_hitl_keeps_text_artifact_and_collects_long_running_parts(): + hitl_part = Part(data=ParseDict({"id": "call-1", "name": "dangerous_tool"}, Value())) + hitl_part.metadata.update( + { + get_kagent_metadata_key("type"): "function_call", + get_kagent_metadata_key("is_long_running"): True, + } + ) + event = TaskArtifactUpdateEvent( + task_id="task-1", + context_id="context-1", + last_chunk=True, + artifact=Artifact( + artifact_id="artifact-1", + parts=[Part(text="please confirm"), hitl_part], + ), + ) + hitl_parts: list[Part] = [] + + kept = _split_hitl_artifact_parts(event, hitl_parts) + + assert kept is event + assert [part.text for part in kept.artifact.parts] == ["please confirm"] + assert len(hitl_parts) == 1 + assert hitl_parts[0].HasField("data") + assert hitl_parts[0].data == hitl_part.data + + +def test_split_hitl_drops_artifact_when_only_long_running_parts_remain(): + hitl_part = Part(data=ParseDict({"id": "call-1", "name": "adk_request_confirmation"}, Value())) + hitl_part.metadata.update({get_kagent_metadata_key("is_long_running"): True}) + event = TaskArtifactUpdateEvent( + task_id="task-1", + context_id="context-1", + last_chunk=True, + artifact=Artifact(artifact_id="artifact-1", parts=[hitl_part]), + ) + hitl_parts: list[Part] = [] + + assert _split_hitl_artifact_parts(event, hitl_parts) is None + assert hitl_parts == [hitl_part] diff --git a/python/packages/kagent-core/src/kagent/core/a2a/__init__.py b/python/packages/kagent-core/src/kagent/core/a2a/__init__.py index 110b8f1cf5..62bae75e8c 100644 --- a/python/packages/kagent-core/src/kagent/core/a2a/__init__.py +++ b/python/packages/kagent-core/src/kagent/core/a2a/__init__.py @@ -30,7 +30,6 @@ ) from ._request_size import A2ARequestSizeLimitMiddleware from ._requests import KAgentRequestContextBuilder -from ._task_result_aggregator import TaskResultAggregator from ._task_store import KAgentTaskStore from ._time import now_timestamp @@ -51,7 +50,6 @@ "A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE", "A2A_DATA_PART_METADATA_TYPE_CODE_EXECUTION_RESULT", "A2A_DATA_PART_METADATA_TYPE_EXECUTABLE_CODE", - "TaskResultAggregator", # HITL constants "KAGENT_HITL_DECISION_TYPE_KEY", "KAGENT_HITL_DECISION_TYPE_APPROVE", diff --git a/python/packages/kagent-core/src/kagent/core/a2a/_task_result_aggregator.py b/python/packages/kagent-core/src/kagent/core/a2a/_task_result_aggregator.py deleted file mode 100644 index b1aa209a33..0000000000 --- a/python/packages/kagent-core/src/kagent/core/a2a/_task_result_aggregator.py +++ /dev/null @@ -1,49 +0,0 @@ -from a2a.server.events import Event -from a2a.types import Message, TaskState, TaskStatusUpdateEvent - - -class TaskResultAggregator: - """Aggregates the task status updates and provides the final task state.""" - - def __init__(self): - self._task_state = TaskState.TASK_STATE_WORKING - self._task_status_message = None - - def process_event(self, event: Event): - """Process an event from the agent run and detect signals about the task status. - Priority of task state: - - failed - - auth_required - - input_required - - working - """ - if isinstance(event, TaskStatusUpdateEvent): - if event.status.state == TaskState.TASK_STATE_FAILED: - self._task_state = TaskState.TASK_STATE_FAILED - self._task_status_message = event.status.message - elif ( - event.status.state == TaskState.TASK_STATE_AUTH_REQUIRED - and self._task_state != TaskState.TASK_STATE_FAILED - ): - self._task_state = TaskState.TASK_STATE_AUTH_REQUIRED - self._task_status_message = event.status.message - elif event.status.state == TaskState.TASK_STATE_INPUT_REQUIRED and self._task_state not in ( - TaskState.TASK_STATE_FAILED, - TaskState.TASK_STATE_AUTH_REQUIRED, - ): - self._task_state = TaskState.TASK_STATE_INPUT_REQUIRED - self._task_status_message = event.status.message - # final state is already recorded and make sure the intermediate state is - # always working because other state may terminate the event aggregation - # in a2a request handler - elif self._task_state == TaskState.TASK_STATE_WORKING: - self._task_status_message = event.status.message - event.status.state = TaskState.TASK_STATE_WORKING - - @property - def task_state(self) -> TaskState: - return self._task_state - - @property - def task_status_message(self) -> Message | None: - return self._task_status_message diff --git a/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py b/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py index f6492e0c9e..16da3d2489 100644 --- a/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py +++ b/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py @@ -4,13 +4,11 @@ import httpx from a2a.server.tasks import TaskStore -from a2a.types import ListTasksRequest, ListTasksResponse, Message, Task +from a2a.types import ListTasksRequest, ListTasksResponse, Task from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE from google.protobuf.json_format import MessageToDict, ParseDict from typing_extensions import override -from kagent.core.a2a import read_metadata_value - logger = logging.getLogger(__name__) @@ -29,23 +27,10 @@ def __init__(self, client: httpx.AsyncClient): # Event-based sync: track pending save operations self._save_events: dict[str, asyncio.Event] = {} - def _is_partial_event(self, item: Message) -> bool: - """Check if a history item is a partial ADK streaming event.""" - metadata = MessageToDict(item.metadata) if item.metadata else {} - return read_metadata_value(metadata, "adk_partial") is True - - def _clean_partial_events(self, history: list[Message]) -> list[Message]: - """Remove partial streaming events from history.""" - return [item for item in history if not self._is_partial_event(item)] - @override async def save(self, task: Task, context=None) -> None: """Save a task to KAgent. - Skips saving if the current event is a partial streaming chunk. - The adk_partial flag is set on event.metadata by AgentExecutor and - gets copied to task.metadata by TaskManager. - Args: task: The task to save context: Server call context (unused, for a2a-sdk 0.3+ compatibility) @@ -53,13 +38,6 @@ async def save(self, task: Task, context=None) -> None: Raises: httpx.HTTPStatusError: If the API request fails """ - # Clean any partial events from history before saving - history = list(task.history or []) - clean_history = self._clean_partial_events(history) - if len(clean_history) != len(history): - del task.history[:] - task.history.extend(clean_history) - response = await self.client.post( "/api/tasks", json=MessageToDict(task), diff --git a/python/packages/kagent-crewai/src/kagent/crewai/_listeners.py b/python/packages/kagent-crewai/src/kagent/crewai/_listeners.py index 9b6ba2072d..a98ad8deae 100644 --- a/python/packages/kagent-crewai/src/kagent/crewai/_listeners.py +++ b/python/packages/kagent-crewai/src/kagent/crewai/_listeners.py @@ -1,13 +1,16 @@ import asyncio +import json import uuid from typing import Any from a2a.server.agent_execution.context import RequestContext from a2a.server.events.event_queue import EventQueue from a2a.types import ( + Artifact, Message, Part, Role, + TaskArtifactUpdateEvent, TaskState, TaskStatus, TaskStatusUpdateEvent, @@ -35,6 +38,38 @@ ) +def _agent_tool_name(agent_name: str) -> str: + """Encode an agent name so the UI renders it via AgentCallDisplay (__NS__).""" + safe = agent_name.replace(" ", "_") + if "/" in safe: + return safe.replace("/", "__NS__") + return f"{safe}__NS__agent" + + +def _agent_display_name(agent: Any) -> str: + return getattr(agent, "role", None) or getattr(agent, "name", None) or str(getattr(agent, "id", "agent")) + + +def _agent_call_id(agent: Any, task: Any) -> str: + agent_id = str(getattr(agent, "id", None) or _agent_display_name(agent)) + task_id = getattr(task, "id", None) if task is not None else None + return f"{agent_id}:{task_id}" if task_id else agent_id + + +def _as_args(raw: Any) -> dict: + if isinstance(raw, dict): + return raw + if isinstance(raw, str): + try: + parsed = json.loads(raw) + return parsed if isinstance(parsed, dict) else {"raw": raw} + except (json.JSONDecodeError, TypeError): + return {"raw": raw} + if raw is None: + return {} + return {"raw": str(raw)} + + class A2ACrewAIListener(BaseEventListener): def __init__( self, @@ -42,202 +77,159 @@ def __init__( event_queue: EventQueue, app_name: str, ): - super().__init__() + # Handlers close over self; fields must exist before super() registers them. self.context = context self.event_queue = event_queue self.app_name = app_name self.loop = asyncio.get_running_loop() + # Stack of in-flight tool call IDs keyed by a stable tool invocation fingerprint. + self._tool_call_ids: dict[tuple[str, str, str, str], list[str]] = {} + super().__init__() def _enqueue_event(self, event: Any): asyncio.run_coroutine_threadsafe(self.event_queue.enqueue_event(event), self.loop) + def _base_metadata(self) -> dict[str, str]: + return { + get_kagent_metadata_key("app_name"): self.app_name, + get_kagent_metadata_key("session_id"): self.context.context_id or "", + } + + def _enqueue_parts(self, parts: list[Part], *, event_type: str | None = None): + metadata = self._base_metadata() + if event_type: + metadata[get_kagent_metadata_key("event_type")] = event_type + self._enqueue_event( + TaskArtifactUpdateEvent( + task_id=self.context.task_id, + context_id=self.context.context_id, + last_chunk=True, + artifact=Artifact(artifact_id=str(uuid.uuid4()), parts=parts, metadata=metadata), + metadata=metadata, + ) + ) + + def _enqueue_status(self, text: str): + """Emit WORKING status for progress tracking (not chat transcript).""" + metadata = self._base_metadata() + self._enqueue_event( + TaskStatusUpdateEvent( + task_id=self.context.task_id, + context_id=self.context.context_id, + status=TaskStatus( + state=TaskState.TASK_STATE_WORKING, + message=Message( + message_id=str(uuid.uuid4()), + role=Role.ROLE_AGENT, + parts=[Part(text=text)], + ), + timestamp=now_timestamp(), + ), + metadata=metadata, + ) + ) + + def _enqueue_function_part(self, data: dict, part_type: str, *, event_type: str): + self._enqueue_parts( + [ + Part( + data=ParseDict(data, Value()), + metadata={get_kagent_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY): part_type}, + ) + ], + event_type=event_type, + ) + + def _tool_key(self, event: ToolUsageStartedEvent | ToolUsageFinishedEvent) -> tuple[str, str, str, str]: + return ( + event.tool_name or "", + event.agent_id or "", + event.task_id or "", + json.dumps(_as_args(event.tool_args), sort_keys=True, default=str), + ) + + def _begin_tool_call(self, event: ToolUsageStartedEvent) -> str: + call_id = str(uuid.uuid4()) + self._tool_call_ids.setdefault(self._tool_key(event), []).append(call_id) + return call_id + + def _end_tool_call(self, event: ToolUsageFinishedEvent) -> str: + key = self._tool_key(event) + stack = self._tool_call_ids.get(key) + if stack: + call_id = stack.pop(0) + if not stack: + del self._tool_call_ids[key] + return call_id + return str(uuid.uuid4()) + def setup_listeners(self, crewai_event_bus): @crewai_event_bus.on(TaskStartedEvent) def on_task_started(source: Any, event: TaskStartedEvent): - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[Part(text=f"Task started: {event.task.name}")], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) - ) + task_name = getattr(event.task, "name", None) or "task" + self._enqueue_status(f"Task started: {task_name}") @crewai_event_bus.on(TaskCompletedEvent) def on_task_completed(source: Any, event: TaskCompletedEvent): - if event.output: - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[Part(text=f"Task completed: {event.task.name}\n")], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) - ) + task_name = getattr(event.task, "name", None) or "task" + self._enqueue_status(f"Task completed: {task_name}") + + @crewai_event_bus.on(MethodExecutionStartedEvent) + def on_method_execution_started(source: Any, event: MethodExecutionStartedEvent): + self._enqueue_status(f"Flow {event.flow_name}: {event.method_name} started") + + @crewai_event_bus.on(MethodExecutionFinishedEvent) + def on_method_execution_finished(source: Any, event: MethodExecutionFinishedEvent): + self._enqueue_status(f"Flow {event.flow_name}: {event.method_name} finished") @crewai_event_bus.on(AgentExecutionStartedEvent) def on_agent_execution_started(source: Any, event: AgentExecutionStartedEvent): - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[Part(text=f"Agent {event.agent.id} started working on task: {event.task_prompt}")], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) + agent_name = _agent_display_name(event.agent) + self._enqueue_function_part( + { + "id": _agent_call_id(event.agent, event.task), + "name": _agent_tool_name(agent_name), + "args": {"task": event.task_prompt or ""}, + }, + A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL, + event_type="agent_execution", ) @crewai_event_bus.on(AgentExecutionCompletedEvent) def on_agent_execution_completed(source: Any, event: AgentExecutionCompletedEvent): - if event.output: - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[Part(text=str(event.output))], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) - ) + agent_name = _agent_display_name(event.agent) + self._enqueue_function_part( + { + "id": _agent_call_id(event.agent, event.task), + "name": _agent_tool_name(agent_name), + "response": {"result": event.output if event.output is not None else ""}, + }, + A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, + event_type="agent_execution", + ) @crewai_event_bus.on(ToolUsageStartedEvent) def on_tool_usage_started(source: Any, event: ToolUsageStartedEvent): - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[ - Part( - data=ParseDict( - { - "id": event.tool_class, - "name": event.tool_name, - "args": event.tool_args, - }, - Value(), - ), - metadata={ - get_kagent_metadata_key( - A2A_DATA_PART_METADATA_TYPE_KEY - ): A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL - }, - ) - ], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) + call_id = self._begin_tool_call(event) + self._enqueue_function_part( + { + "id": call_id, + "name": event.tool_name, + "args": _as_args(event.tool_args), + }, + A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL, + event_type="tool_call", ) @crewai_event_bus.on(ToolUsageFinishedEvent) def on_tool_usage_finished(source: Any, event: ToolUsageFinishedEvent): - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[ - Part( - data=ParseDict( - { - "id": event.tool_class, - "name": event.tool_name, - "response": event.output, - }, - Value(), - ), - metadata={ - get_kagent_metadata_key( - A2A_DATA_PART_METADATA_TYPE_KEY - ): A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, - }, - ) - ], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) - ) - - @crewai_event_bus.on(MethodExecutionStartedEvent) - def on_method_execution_started(source: Any, event: MethodExecutionStartedEvent): - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[ - Part(text=f"Method {event.method_name} from flow {event.flow_name} started execution.") - ], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) - ) - - @crewai_event_bus.on(MethodExecutionFinishedEvent) - def on_method_execution_finished(source: Any, event: MethodExecutionFinishedEvent): - self._enqueue_event( - TaskStatusUpdateEvent( - task_id=self.context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[ - Part(text=f"Method {event.method_name} from flow {event.flow_name} finished execution.") - ], - ), - ), - context_id=self.context.context_id, - metadata={"app_name": self.app_name, "session_id": self.context.context_id}, - ) + call_id = self._end_tool_call(event) + self._enqueue_function_part( + { + "id": call_id, + "name": event.tool_name, + "response": {"result": event.output}, + }, + A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, + event_type="tool_output", ) diff --git a/python/packages/kagent-crewai/tests/test_executor.py b/python/packages/kagent-crewai/tests/test_executor.py index 550f9fd2a9..17f17a2d65 100644 --- a/python/packages/kagent-crewai/tests/test_executor.py +++ b/python/packages/kagent-crewai/tests/test_executor.py @@ -3,8 +3,7 @@ import httpx import pytest from a2a.server.agent_execution.context import RequestContext -from a2a.server.events.event_queue import EventQueue -from a2a.types import Message, Part, Role, SendMessageRequest +from a2a.types import Message, Part, Role, SendMessageRequest, TaskArtifactUpdateEvent from google.protobuf.json_format import ParseDict from google.protobuf.struct_pb2 import Value @@ -29,13 +28,30 @@ def _make_crew() -> MagicMock: return crew -async def _run(crew: MagicMock, context: RequestContext) -> None: +class _RecordingEventQueue: + def __init__(self): + self.events = [] + + async def enqueue_event(self, event): + self.events.append(event) + + +async def _run(crew: MagicMock, context: RequestContext) -> list: executor = CrewAIAgentExecutor( crew=crew, app_name="test", http_client=httpx.AsyncClient(), ) - await executor.execute(context, EventQueue()) + event_queue = _RecordingEventQueue() + await executor.execute(context, event_queue) + return event_queue.events + + +def _assert_content_artifact_closes_stream(events: list) -> None: + artifacts = [event for event in events if isinstance(event, TaskArtifactUpdateEvent)] + assert artifacts + assert all(artifact.artifact.parts for artifact in artifacts) + assert artifacts[-1].last_chunk is True @pytest.mark.asyncio @@ -43,9 +59,10 @@ async def test_execute_passes_datapart_data_as_inputs(): crew = _make_crew() context = _request_context(Part(data=ParseDict({"topic": "ai"}, Value()))) - await _run(crew, context) + events = await _run(crew, context) crew.kickoff_async.assert_awaited_once_with(inputs={"topic": "ai"}) + _assert_content_artifact_closes_stream(events) @pytest.mark.asyncio @@ -53,6 +70,7 @@ async def test_execute_falls_back_to_text_input_without_datapart(): crew = _make_crew() context = _request_context(Part(text="hello")) - await _run(crew, context) + events = await _run(crew, context) crew.kickoff_async.assert_awaited_once_with(inputs={"input": "hello"}) + _assert_content_artifact_closes_stream(events) diff --git a/python/packages/kagent-langgraph/src/kagent/langgraph/_converters.py b/python/packages/kagent-langgraph/src/kagent/langgraph/_converters.py index 27e808df95..0ea5be6b09 100644 --- a/python/packages/kagent-langgraph/src/kagent/langgraph/_converters.py +++ b/python/packages/kagent-langgraph/src/kagent/langgraph/_converters.py @@ -9,12 +9,11 @@ from typing import Any from a2a.types import ( + Artifact, Message, Part, Role, - TaskState, - TaskStatus, - TaskStatusUpdateEvent, + TaskArtifactUpdateEvent, ) from google.protobuf.json_format import ParseDict from google.protobuf.struct_pb2 import Value @@ -23,7 +22,6 @@ A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, A2A_DATA_PART_METADATA_TYPE_KEY, get_kagent_metadata_key, - now_timestamp, ) from langchain_core.messages import ( AIMessage, @@ -40,12 +38,12 @@ async def _convert_langgraph_event_to_a2a( context_id: str, app_name: str, sent_message_ids: set[str], -) -> list[TaskStatusUpdateEvent]: +) -> list[TaskArtifactUpdateEvent]: """Convert a LangGraph event to A2A events. Deduplicates messages using sent_message_ids to avoid replaying history. """ - a2a_events: list[TaskStatusUpdateEvent] = [] + a2a_events: list[TaskArtifactUpdateEvent] = [] # LangGraph events have node names as keys, with 'messages' as values # Example: {'agent': {'messages': [AIMessage(...)]}} @@ -98,58 +96,48 @@ async def _convert_langgraph_event_to_a2a( if not a2a_message.parts: continue + metadata = get_rich_event_metadata(app_name=app_name, session_id=context_id) a2a_events.append( - TaskStatusUpdateEvent( + TaskArtifactUpdateEvent( task_id=task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=a2a_message, - ), context_id=context_id, - metadata=get_rich_event_metadata( - app_name=app_name, - session_id=context_id, - ), + last_chunk=True, + artifact=Artifact(artifact_id=str(uuid.uuid4()), parts=a2a_message.parts, metadata=metadata), + metadata=metadata, ) ) elif isinstance(message, ToolMessage): # Handle tool responses if message.content: + metadata = get_rich_event_metadata(app_name=app_name, session_id=context_id) a2a_events.append( - TaskStatusUpdateEvent( + TaskArtifactUpdateEvent( task_id=task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - timestamp=now_timestamp(), - message=Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[ - Part( - data=ParseDict( - { - "id": message.tool_call_id, - "name": message.name, - "response": message.content, - }, - Value(), - ), - metadata={ - get_kagent_metadata_key( - A2A_DATA_PART_METADATA_TYPE_KEY - ): A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, - }, - ) - ], - ), - ), context_id=context_id, - metadata=get_rich_event_metadata( - app_name=app_name, - session_id=context_id, + last_chunk=True, + artifact=Artifact( + artifact_id=str(uuid.uuid4()), + parts=[ + Part( + data=ParseDict( + { + "id": message.tool_call_id, + "name": message.name, + "response": message.content, + }, + Value(), + ), + metadata={ + get_kagent_metadata_key( + A2A_DATA_PART_METADATA_TYPE_KEY + ): A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, + }, + ) + ], + metadata=metadata, ), + metadata=metadata, ) ) diff --git a/python/packages/kagent-langgraph/src/kagent/langgraph/_executor.py b/python/packages/kagent-langgraph/src/kagent/langgraph/_executor.py index 1d2f9e5935..935480a0f9 100644 --- a/python/packages/kagent-langgraph/src/kagent/langgraph/_executor.py +++ b/python/packages/kagent-langgraph/src/kagent/langgraph/_executor.py @@ -19,12 +19,10 @@ from a2a.server.agent_execution.context import RequestContext from a2a.server.events.event_queue import EventQueue from a2a.types import ( - Artifact, Message, Part, Role, Task, - TaskArtifactUpdateEvent, TaskState, TaskStatus, TaskStatusUpdateEvent, @@ -37,7 +35,6 @@ A2A_DATA_PART_METADATA_TYPE_KEY, KAGENT_HITL_DECISION_TYPE_BATCH, KAGENT_HITL_DECISION_TYPE_REJECT, - TaskResultAggregator, extract_ask_user_answers_from_message, extract_batch_decisions_from_message, extract_decision_from_message, @@ -136,8 +133,6 @@ async def _stream_graph_events( event_queue: EventQueue, ) -> None: """Stream LangGraph events and convert them to A2A events.""" - task_result_aggregator = TaskResultAggregator() - # Track final state for interrupt detection final_state: dict[str, Any] | None = None @@ -158,7 +153,6 @@ async def _stream_graph_events( event, context.task_id, context.context_id, self.app_name, sent_message_ids ) for a2a_event in a2a_events: - task_result_aggregator.process_event(a2a_event) await event_queue.enqueue_event(a2a_event) # Check for interrupts after streaming completes @@ -173,50 +167,16 @@ async def _stream_graph_events( # Interrupt detected - input_required event already sent, so return early return - # Final artifacts are already sent through individual event processing - - # publish the task result event - this is final - if ( - task_result_aggregator.task_state == TaskState.TASK_STATE_WORKING - and task_result_aggregator.task_status_message is not None - and task_result_aggregator.task_status_message.parts - ): - # if task is still working properly, publish the artifact update event as - # the final result according to a2a protocol. - await event_queue.enqueue_event( - TaskArtifactUpdateEvent( - task_id=context.task_id, - last_chunk=True, - context_id=context.context_id, - artifact=Artifact( - artifact_id=str(uuid.uuid4()), - parts=task_result_aggregator.task_status_message.parts, - ), - ) - ) - # public the final status update event - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_COMPLETED, - timestamp=now_timestamp(), - ), - context_id=context.context_id, - ) - ) - else: - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=context.task_id, - status=TaskStatus( - state=task_result_aggregator.task_state, - timestamp=now_timestamp(), - message=task_result_aggregator.task_status_message, - ), - context_id=context.context_id, - ) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=context.task_id, + status=TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + timestamp=now_timestamp(), + ), + context_id=context.context_id, ) + ) async def _handle_interrupt( self, diff --git a/python/packages/kagent-openai/src/kagent/openai/_agent_executor.py b/python/packages/kagent-openai/src/kagent/openai/_agent_executor.py index 9ed4160dd1..7aeade6321 100644 --- a/python/packages/kagent-openai/src/kagent/openai/_agent_executor.py +++ b/python/packages/kagent-openai/src/kagent/openai/_agent_executor.py @@ -33,7 +33,7 @@ ) from agents.agent import Agent from agents.run import Runner -from kagent.core.a2a import TaskResultAggregator, get_kagent_metadata_key, now_timestamp +from kagent.core.a2a import get_kagent_metadata_key, now_timestamp from pydantic import BaseModel from ._event_converter import convert_openai_event_to_a2a_events @@ -101,8 +101,8 @@ async def _stream_agent_events( event_queue: EventQueue, ) -> None: """Stream agent execution events and convert them to A2A events.""" - task_result_aggregator = TaskResultAggregator() session_context = SessionContext(session_id=session.session_id) + emitted_text = False try: # Use run_streamed for streaming support @@ -124,18 +124,12 @@ async def _stream_agent_events( ) for a2a_event in a2a_events: - task_result_aggregator.process_event(a2a_event) + if isinstance(a2a_event, TaskArtifactUpdateEvent): + emitted_text = emitted_text or any(part.HasField("text") for part in a2a_event.artifact.parts) await event_queue.enqueue_event(a2a_event) - # Handle final output - if hasattr(result, "final_output") and result.final_output: - final_message = Message( - message_id=str(uuid.uuid4()), - role=Role.ROLE_AGENT, - parts=[Part(text=str(result.final_output))], - ) - - # Publish final artifact + # Some SDK runs expose a final output without a corresponding stream item. + if not emitted_text and hasattr(result, "final_output") and result.final_output: await event_queue.enqueue_event( TaskArtifactUpdateEvent( task_id=context.task_id, @@ -143,62 +137,21 @@ async def _stream_agent_events( context_id=context.context_id, artifact=Artifact( artifact_id=str(uuid.uuid4()), - parts=final_message.parts, + parts=[Part(text=str(result.final_output))], ), ) ) - # Publish completion status - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_COMPLETED, - timestamp=now_timestamp(), - ), - context_id=context.context_id, - ) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=context.task_id, + status=TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + timestamp=now_timestamp(), + ), + context_id=context.context_id, ) - else: - # No output - publish based on aggregator state - if ( - task_result_aggregator.task_state == TaskState.TASK_STATE_WORKING - and task_result_aggregator.task_status_message is not None - and task_result_aggregator.task_status_message.parts - ): - await event_queue.enqueue_event( - TaskArtifactUpdateEvent( - task_id=context.task_id, - last_chunk=True, - context_id=context.context_id, - artifact=Artifact( - artifact_id=str(uuid.uuid4()), - parts=task_result_aggregator.task_status_message.parts, - ), - ) - ) - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=context.task_id, - status=TaskStatus( - state=TaskState.TASK_STATE_COMPLETED, - timestamp=now_timestamp(), - ), - context_id=context.context_id, - ) - ) - else: - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=context.task_id, - status=TaskStatus( - state=task_result_aggregator.task_state, - timestamp=now_timestamp(), - message=task_result_aggregator.task_status_message, - ), - context_id=context.context_id, - ) - ) + ) except Exception as e: logger.error(f"Error during agent execution: {e}", exc_info=True) diff --git a/python/packages/kagent-openai/src/kagent/openai/_event_converter.py b/python/packages/kagent-openai/src/kagent/openai/_event_converter.py index e1a9bb3c2e..7d413b3414 100644 --- a/python/packages/kagent-openai/src/kagent/openai/_event_converter.py +++ b/python/packages/kagent-openai/src/kagent/openai/_event_converter.py @@ -11,16 +11,14 @@ from a2a.server.events import Event as A2AEvent from a2a.types import ( + Artifact, Message, Role, - TaskState, - TaskStatus, - TaskStatusUpdateEvent, + TaskArtifactUpdateEvent, ) from a2a.types import Part as A2APart -from agents.items import MessageOutputItem, ToolCallItem, ToolCallOutputItem +from agents.items import HandoffCallItem, HandoffOutputItem, MessageOutputItem, ToolCallItem, ToolCallOutputItem from agents.stream_events import ( - AgentUpdatedStreamEvent, RawResponsesStreamEvent, RunItemStreamEvent, StreamEvent, @@ -32,12 +30,25 @@ A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, A2A_DATA_PART_METADATA_TYPE_KEY, get_kagent_metadata_key, - now_timestamp, ) logger = logging.getLogger(__name__) +def _artifact_event(message: Message, task_id: str, context_id: str) -> TaskArtifactUpdateEvent: + return TaskArtifactUpdateEvent( + task_id=task_id, + context_id=context_id, + last_chunk=True, + artifact=Artifact( + artifact_id=str(uuid.uuid4()), + parts=message.parts, + metadata=message.metadata, + ), + metadata=message.metadata, + ) + + def convert_openai_event_to_a2a_events( event: StreamEvent, task_id: str, @@ -67,10 +78,6 @@ def convert_openai_event_to_a2a_events( # These are low-level events - can be logged but not converted logger.debug(f"Raw response event: {event.data}") - # Handle AgentUpdatedStreamEvent (agent handoffs) - elif isinstance(event, AgentUpdatedStreamEvent): - a2a_events.extend(_convert_agent_updated_event(event, task_id, context_id, app_name)) - # Other event types else: logger.debug(f"Unhandled event type: {type(event).__name__}") @@ -111,6 +118,14 @@ def _convert_run_item_event( elif isinstance(event.item, ToolCallOutputItem): return _convert_tool_output(event.item, task_id, context_id, app_name) + # Handle handoff calls (map to subagent-style function_call for the UI) + elif isinstance(event.item, HandoffCallItem): + return _convert_handoff_call(event.item, task_id, context_id, app_name) + + # Handle handoff outputs (map to subagent-style function_response) + elif isinstance(event.item, HandoffOutputItem): + return _convert_handoff_output(event.item, task_id, context_id, app_name) + # Other item types else: logger.debug(f"Unhandled run item type: {type(event.item).__name__}") @@ -161,20 +176,7 @@ def _convert_message_output( }, ) - status_event = TaskStatusUpdateEvent( - task_id=task_id, - context_id=context_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - message=message, - timestamp=now_timestamp(), - ), - metadata={ - get_kagent_metadata_key("app_name"): app_name, - }, - ) - - return [status_event] + return [_artifact_event(message, task_id, context_id)] def _convert_tool_call( @@ -237,20 +239,7 @@ def _convert_tool_call( }, ) - status_event = TaskStatusUpdateEvent( - task_id=task_id, - context_id=context_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - message=message, - timestamp=now_timestamp(), - ), - metadata={ - get_kagent_metadata_key("app_name"): app_name, - }, - ) - - return [status_event] + return [_artifact_event(message, task_id, context_id)] def _convert_tool_output( @@ -299,43 +288,61 @@ def _convert_tool_output( }, ) - status_event = TaskStatusUpdateEvent( - task_id=task_id, - context_id=context_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - message=message, - timestamp=now_timestamp(), - ), - metadata={ - get_kagent_metadata_key("app_name"): app_name, - }, - ) + return [_artifact_event(message, task_id, context_id)] + + +def _agent_tool_name(agent_name: str) -> str: + """Encode an agent name so the UI renders it via AgentCallDisplay (__NS__).""" + if "/" in agent_name: + return agent_name.replace("/", "__NS__") + return f"{agent_name}__NS__agent" + + +def _parse_tool_arguments(raw_arguments: object) -> dict: + if isinstance(raw_arguments, dict): + return raw_arguments + if isinstance(raw_arguments, str): + try: + parsed = json.loads(raw_arguments) + return parsed if isinstance(parsed, dict) else {"raw": raw_arguments} + except (json.JSONDecodeError, TypeError): + logger.warning(f"Failed to parse arguments: {raw_arguments}") + return {"raw": raw_arguments} + if raw_arguments is None: + return {} + return {"raw": str(raw_arguments)} - return [status_event] +def _handoff_target_from_call(raw_call: object) -> str: + """Best-effort target agent name from a transfer_to_* handoff tool call.""" + tool_name = getattr(raw_call, "name", None) or "unknown" + if tool_name.startswith("transfer_to_"): + return tool_name.removeprefix("transfer_to_") + return tool_name -def _convert_agent_updated_event( - event: AgentUpdatedStreamEvent, + +def _convert_handoff_call( + item: HandoffCallItem, task_id: str, context_id: str, app_name: str, ) -> list[A2AEvent]: - """Convert an agent updated event (handoff) to A2A event. - - This is converted to a function_call event so the frontend renders it - using the AgentCallDisplay component. This is ideal if there are multiple handoffs. - """ - agent_name = event.new_agent.name - if "/" in agent_name: - tool_name = agent_name.replace("/", "__NS__") - else: - tool_name = f"{agent_name}__NS__agent" + """Convert a handoff request to a subagent-style function_call A2A event.""" + raw_call = item.raw_item + call_id = ( + raw_call.call_id + if hasattr(raw_call, "call_id") and raw_call.call_id + else (raw_call.id if hasattr(raw_call, "id") and raw_call.id else str(uuid.uuid4())) + ) + agent_name = _handoff_target_from_call(raw_call) + tool_arguments = _parse_tool_arguments(getattr(raw_call, "arguments", None)) + if "target_agent" not in tool_arguments: + tool_arguments = {**tool_arguments, "target_agent": agent_name} function_data = { - "id": str(uuid.uuid4()), - "name": tool_name, - "args": {"target_agent": agent_name}, + "id": call_id, + "name": _agent_tool_name(agent_name), + "args": tool_arguments, } message = Message( @@ -356,17 +363,49 @@ def _convert_agent_updated_event( }, ) - status_event = TaskStatusUpdateEvent( - task_id=task_id, - context_id=context_id, - status=TaskStatus( - state=TaskState.TASK_STATE_WORKING, - message=message, - timestamp=now_timestamp(), - ), + return [_artifact_event(message, task_id, context_id)] + + +def _convert_handoff_output( + item: HandoffOutputItem, + task_id: str, + context_id: str, + app_name: str, +) -> list[A2AEvent]: + """Convert a handoff output to a subagent-style function_response A2A event.""" + raw_output = item.raw_item + if isinstance(raw_output, dict): + call_id = raw_output.get("call_id") or str(uuid.uuid4()) + result = raw_output.get("output", "") + else: + call_id = getattr(raw_output, "call_id", None) or str(uuid.uuid4()) + result = getattr(raw_output, "output", "") + + agent_name = item.target_agent.name if item.target_agent else "unknown" + function_data = { + "id": call_id, + "name": _agent_tool_name(agent_name), + "response": {"result": result}, + } + + message = Message( + message_id=str(uuid.uuid4()), + role=Role.ROLE_AGENT, + parts=[ + A2APart( + data=ParseDict(function_data, Value()), + metadata={ + get_kagent_metadata_key( + A2A_DATA_PART_METADATA_TYPE_KEY + ): A2A_DATA_PART_METADATA_TYPE_FUNCTION_RESPONSE, + }, + ) + ], metadata={ get_kagent_metadata_key("app_name"): app_name, + get_kagent_metadata_key("event_type"): "agent_handoff_output", + get_kagent_metadata_key("new_agent_name"): agent_name, }, ) - return [status_event] + return [_artifact_event(message, task_id, context_id)] diff --git a/ui/playwright/mocks/server.mjs b/ui/playwright/mocks/server.mjs index 05ac173fd0..61323c8637 100644 --- a/ui/playwright/mocks/server.mjs +++ b/ui/playwright/mocks/server.mjs @@ -59,23 +59,30 @@ const AGENT_REPLY = process.env.E2E_AGENT_REPLY ?? "Hello from the agent"; const frame = (event) => `data: ${JSON.stringify({ jsonrpc: "2.0", id: "1", result: event })}\n\n`; function chatStream(contextId, taskId) { + const artifact = { + artifactUpdate: { + taskId, + contextId, + artifact: { + artifactId: "e2e-agent-reply", + name: "", + description: "", + parts: [{ text: AGENT_REPLY }], + }, + append: false, + lastChunk: true, + }, + }; const completed = { statusUpdate: { taskId, contextId, status: { state: "TASK_STATE_COMPLETED", - message: { - messageId: "e2e-agent-reply", - role: "ROLE_AGENT", - parts: [{ text: AGENT_REPLY }], - contextId, - taskId, - }, }, }, }; - return frame(completed) + "data: [DONE]\n\n"; + return frame(artifact) + frame(completed) + "data: [DONE]\n\n"; } async function handleChat(req, res) { diff --git a/ui/src/components/chat/ChatInterface.tsx b/ui/src/components/chat/ChatInterface.tsx index 4164859a88..f060f8c661 100644 --- a/ui/src/components/chat/ChatInterface.tsx +++ b/ui/src/components/chat/ChatInterface.tsx @@ -30,7 +30,18 @@ import { getUiRuntimeConfig } from "@/app/actions/config"; import { DEFAULT_STREAM_TIMEOUT_MS } from "@/lib/constants"; import { toast } from "sonner"; import { useRouter } from "next/navigation"; -import { createMessageHandlers, extractMessagesFromTasks, extractApprovalMessagesFromTasks, extractTokenStatsFromTasks, createMessage, ADKMetadata, ProcessedToolCallData } from "@/lib/messageHandlers"; +import { + createMessageHandlers, + extractMessagesFromTasks, + extractApprovalMessagesFromTasks, + extractTokenStatsFromTasks, + collectTerminalTaskIds, + collectTaskTokenStats, + isFinishedAssistantReply, + createMessage, + ADKMetadata, + ProcessedToolCallData, +} from "@/lib/messageHandlers"; import { kagentA2AClient } from "@/lib/a2aClient"; import { formatA2AClientError } from "@/lib/a2aErrors"; import { useChatRunInSandbox, useChatSubstrateSandbox, useCurrentChatAgent } from "@/components/chat/ChatAgentContext"; @@ -64,6 +75,7 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se const [currentInputMessage, setCurrentInputMessage] = useState(""); const [chatStatus, setChatStatus] = useState("ready"); + const [statusMessage, setStatusMessage] = useState(undefined); const [session, setSession] = useState(selectedSession || null); const [shareReadOnly, setShareReadOnly] = useState(false); @@ -78,6 +90,9 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se const isCreatingSessionRef = useRef(false); const [isFirstMessage, setIsFirstMessage] = useState(!sessionId); const [sessionStats, setSessionStats] = useState({ total: 0, prompt: 0, completion: 0 }); + // Finished-reply chrome is derived from terminal task status, not message metadata. + const [terminalTaskIds, setTerminalTaskIds] = useState>(() => new Set()); + const [taskTokenStats, setTaskTokenStats] = useState>(() => new Map()); // Mutable ref so pendingTurnStats survives re-renders between A2A stream events const pendingTurnStatsRef = useRef(undefined); const [pendingDecisions, setPendingDecisions] = useState>({}); @@ -171,22 +186,42 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se // Shared call_id -> is_error lookup so each group summary is O(group size). const toolResultsByCallId = useMemo(() => buildToolCallResultsIndex(allMessages), [allMessages]); + const onTerminalTask = useCallback((taskId: string, tokenStats?: TokenStats) => { + setTerminalTaskIds(prev => { + if (prev.has(taskId)) return prev; + const next = new Set(prev); + next.add(taskId); + return next; + }); + if (tokenStats) { + setTaskTokenStats(prev => { + const next = new Map(prev); + next.set(taskId, tokenStats); + return next; + }); + } + }, []); + const { handleMessageEvent } = useMemo(() => createMessageHandlers({ setMessages: setStreamingMessages, setIsStreaming, setStreamingContent, setChatStatus, + setStatusMessage, setSessionStats, pendingTurnStats: pendingTurnStatsRef, + onTerminalTask, agentContext: { namespace: selectedNamespace, agentName: selectedAgentName } - }), [selectedNamespace, selectedAgentName]); + }), [selectedNamespace, selectedAgentName, onTerminalTask]); useEffect(() => { async function initializeChat() { setSessionStats({ total: 0, prompt: 0, completion: 0 }); + setTerminalTaskIds(new Set()); + setTaskTokenStats(new Map()); setStreamingMessages([]); setPendingDecisions({}); pendingDecisionsRef.current = {}; @@ -240,13 +275,19 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se if (!messagesResponse.data || messagesResponse?.data?.length === 0) { setStoredMessages([]); setSessionStats({ total: 0, prompt: 0, completion: 0 }); + setTerminalTaskIds(new Set()); + setTaskTokenStats(new Map()); } else { const extractedMessages = extractMessagesFromTasks(messagesResponse.data); setSessionStats(extractTokenStatsFromTasks(messagesResponse.data)); + setTerminalTaskIds(collectTerminalTaskIds(messagesResponse.data)); + setTaskTokenStats(collectTaskTokenStats(messagesResponse.data)); - // Resolved approvals are already inline in extractedMessages (with - // approved/rejected badges). Only pending approvals need appending. + // Artifact order drives the reloaded transcript. Resolved historical + // approvals are included only when extractMessagesFromTasks can + // anchor them to a matching artifact call/response; append the + // current pending interaction after the assembled output. const { messages: pendingApprovalMessages, hasPendingApproval } = extractApprovalMessagesFromTasks(messagesResponse.data); setStoredMessages( @@ -338,6 +379,7 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se setCurrentInputMessage(""); setChatStatus("thinking"); + setStatusMessage(undefined); setStoredMessages(prev => [...prev, ...streamingMessages]); setStreamingMessages([]); setStreamingContent(""); // Reset streaming content for new message @@ -555,6 +597,8 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se : extractedMessages ); setSessionStats(extractTokenStatsFromTasks(latest.data)); + setTerminalTaskIds(collectTerminalTaskIds(latest.data)); + setTaskTokenStats(collectTaskTokenStats(latest.data)); setStreamingMessages([]); if (hasPendingApproval) { setChatStatus("input_required"); @@ -1013,6 +1057,8 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se message={message} allMessages={allMessages} agentContext={agentContext} + showReplyActions={isFinishedAssistantReply(message, allMessages, terminalTaskIds)} + replyTokenStats={message.taskId ? taskTokenStats.get(message.taskId) : undefined} onApprove={shareReadOnly ? undefined : handleApprove} onReject={shareReadOnly ? undefined : handleReject} onAskUserSubmit={shareReadOnly ? undefined : handleAskUserSubmit} @@ -1108,7 +1154,7 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se ) : ( <>
- +
{sessionStats.total > 0 && } {(session?.id ?? sessionId) && !shareToken && } diff --git a/ui/src/components/chat/ChatMessage.tsx b/ui/src/components/chat/ChatMessage.tsx index 60f58d3ae9..41c1fa35ca 100644 --- a/ui/src/components/chat/ChatMessage.tsx +++ b/ui/src/components/chat/ChatMessage.tsx @@ -21,6 +21,9 @@ interface ChatMessageProps { namespace: string; agentName: string; }; + /** Derived from terminal task status + last text for that task */ + showReplyActions?: boolean; + replyTokenStats?: TokenStats; onApprove?: (toolCallId: string) => void; onReject?: (toolCallId: string, reason?: string) => void; onAskUserSubmit?: (answers: Array<{ answer: string[] }>) => void; @@ -29,7 +32,19 @@ interface ChatMessageProps { onMcpAppSendMessage?: (text: string) => Promise; } -export default function ChatMessage({ message, allMessages, agentContext, onApprove, onReject, onAskUserSubmit, pendingDecisions, getMcpAppForTool, onMcpAppSendMessage }: ChatMessageProps) { +export default function ChatMessage({ + message, + allMessages, + agentContext, + showReplyActions = false, + replyTokenStats, + onApprove, + onReject, + onAskUserSubmit, + pendingDecisions, + getMcpAppForTool, + onMcpAppSendMessage, +}: ChatMessageProps) { const [feedbackDialogOpen, setFeedbackDialogOpen] = useState(false); const [isPositiveFeedback, setIsPositiveFeedback] = useState(true); @@ -38,7 +53,6 @@ export default function ChatMessage({ message, allMessages, agentContext, onAppr const content = message.parts?.filter(isTextPart).map((part) => part.content.value).join("") || ""; const source = isUserRole(message.role) ? "user" : "assistant"; - const tokenStats = (message.metadata as Record | undefined)?.tokenStats as TokenStats | undefined; const messageId = message.messageId; // Extract agent name from metadata for display @@ -177,9 +191,9 @@ export default function ChatMessage({ message, allMessages, agentContext, onAppr of its own, which is what long tool-call ids and URLs need. */} - {source !== "user" && ( + {source !== "user" && showReplyActions && (
- {tokenStats && } + {replyTokenStats && } {messageId !== undefined && ( <>