From 72996111b7c4aedc94a144d166bb9a36c2899855 Mon Sep 17 00:00:00 2001 From: trangevi Date: Thu, 2 Apr 2026 10:43:57 -0700 Subject: [PATCH 1/2] specify supported protocols when generating agent.yaml Signed-off-by: trangevi --- .../azure.ai.agents/internal/cmd/init.go | 4 + .../internal/cmd/init_from_code.go | 123 +++++++++++++++++- .../internal/cmd/init_from_code_test.go | 102 +++++++++++++++ 3 files changed, 223 insertions(+), 6 deletions(-) diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init.go index 2aec48d32ae..4ab98fd8e29 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init.go @@ -46,6 +46,7 @@ type initFlags struct { manifestPointer string src string env string + protocols []string } // AiProjectResourceConfig represents the configuration for an AI project resource @@ -354,6 +355,9 @@ func newInitCommand(rootFlags *rootFlagsDefinition) *cobra.Command { cmd.Flags().StringVarP(&flags.env, "environment", "e", "", "The name of the azd environment to use.") + cmd.Flags().StringSliceVar(&flags.protocols, "protocol", nil, + "Protocols supported by the agent (e.g., 'responses', 'invocations'). Can be specified multiple times.") + return cmd } diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go index 232f9af6842..33b18ab69cc 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go @@ -442,6 +442,12 @@ func (a *InitFromCodeAction) createDefinitionFromLocalAgent(ctx context.Context) // TODO: Prompt user for agent kind agentKind := agent_yaml.AgentKindHosted + // Prompt user for supported protocols + protocols, err := promptProtocols(ctx, a.azdClient, a.flags.NoPrompt, a.flags.protocols) + if err != nil { + return nil, err + } + // Ask user how they want to configure a model modelConfigChoices := []*azdext.SelectChoice{ {Label: "Deploy a new model from the catalog", Value: "new"}, @@ -551,12 +557,7 @@ func (a *InitFromCodeAction) createDefinitionFromLocalAgent(ctx context.Context) Name: agentName, Kind: agentKind, }, - Protocols: []agent_yaml.ProtocolVersionRecord{ - { - Protocol: "responses", - Version: "v1", - }, - }, + Protocols: protocols, EnvironmentVariables: &[]agent_yaml.EnvironmentVariable{ { Name: "AZURE_OPENAI_ENDPOINT", @@ -787,3 +788,113 @@ func (a *InitFromCodeAction) addToProject(ctx context.Context, targetDir string, fmt.Printf("\nAdded your agent as a service entry named '%s' under the file azure.yaml.\n", agentName) return nil } + +// protocolInfo pairs a protocol name with the default version used when generating agent.yaml. +type protocolInfo struct { + Name string + Version string +} + +// knownProtocols lists the protocols offered during init, in display order. +var knownProtocols = []protocolInfo{ + {Name: "responses", Version: "v1"}, + {Name: "invocations", Version: "v0.0.1"}, +} + +// promptProtocols asks the user which protocols their agent supports. +// When flagProtocols is non-empty the prompt is skipped and those values are used directly. +// When noPrompt is true and no flag values are provided, defaults to [responses/v1]. +func promptProtocols( + ctx context.Context, + azdClient *azdext.AzdClient, + noPrompt bool, + flagProtocols []string, +) ([]agent_yaml.ProtocolVersionRecord, error) { + // Build a lookup from protocol name → version for known protocols. + versionOf := make(map[string]string, len(knownProtocols)) + for _, p := range knownProtocols { + versionOf[p.Name] = p.Version + } + + // If explicit flag values were provided, use them directly. + if len(flagProtocols) > 0 { + records := make([]agent_yaml.ProtocolVersionRecord, 0, len(flagProtocols)) + for _, name := range flagProtocols { + version, ok := versionOf[name] + if !ok { + return nil, exterrors.Validation( + exterrors.CodeInvalidAgentManifest, + fmt.Sprintf("unknown protocol %q; supported values: %s", + name, knownProtocolNames()), + fmt.Sprintf("Use one of the supported protocol values: %s", knownProtocolNames()), + ) + } + records = append(records, agent_yaml.ProtocolVersionRecord{ + Protocol: name, + Version: version, + }) + } + return records, nil + } + + // Non-interactive mode: default to responses. + if noPrompt { + return []agent_yaml.ProtocolVersionRecord{ + {Protocol: "responses", Version: "v1"}, + }, nil + } + + // Build multi-select choices; "responses" is pre-selected. + choices := make([]*azdext.MultiSelectChoice, 0, len(knownProtocols)) + for _, p := range knownProtocols { + choices = append(choices, &azdext.MultiSelectChoice{ + Value: p.Name, + Label: p.Name, + Selected: p.Name == "responses", + }) + } + + resp, err := azdClient.Prompt().MultiSelect(ctx, &azdext.MultiSelectRequest{ + Options: &azdext.MultiSelectOptions{ + Message: "Which protocols does your agent support?", + Choices: choices, + Hint: "Use arrow keys to move, space to toggle, enter to confirm", + }, + }) + if err != nil { + if exterrors.IsCancellation(err) { + return nil, exterrors.Cancelled("protocol selection was cancelled") + } + return nil, fmt.Errorf("failed to prompt for protocols: %w", err) + } + + // Collect selected protocols. + var records []agent_yaml.ProtocolVersionRecord + for _, choice := range resp.Values { + if choice.Selected { + records = append(records, agent_yaml.ProtocolVersionRecord{ + Protocol: choice.Value, + Version: versionOf[choice.Value], + }) + } + } + + if len(records) == 0 { + return nil, exterrors.Validation( + exterrors.CodeInvalidAgentManifest, + "at least one protocol must be selected", + "Select at least one protocol for your agent.", + ) + } + + return records, nil +} + +// knownProtocolNames returns a comma-separated list of known protocol names. +func knownProtocolNames() string { + names := make([]string, 0, len(knownProtocols)) + for _, p := range knownProtocols { + names = append(names, p.Name) + } + return strings.Join(names, ", ") +} diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go index 06b3a1e68ac..6af3eb50aa3 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go @@ -580,3 +580,105 @@ func containsAll(s string, substrings ...string) bool { } return true } + +func TestPromptProtocols_FlagValues(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + flagProtocols []string + wantProtocols []agent_yaml.ProtocolVersionRecord + wantErr bool + wantErrContain string + }{ + { + name: "responses only", + flagProtocols: []string{"responses"}, + wantProtocols: []agent_yaml.ProtocolVersionRecord{ + {Protocol: "responses", Version: "v1"}, + }, + }, + { + name: "invocations only", + flagProtocols: []string{"invocations"}, + wantProtocols: []agent_yaml.ProtocolVersionRecord{ + {Protocol: "invocations", Version: "v0.0.1"}, + }, + }, + { + name: "both protocols", + flagProtocols: []string{"responses", "invocations"}, + wantProtocols: []agent_yaml.ProtocolVersionRecord{ + {Protocol: "responses", Version: "v1"}, + {Protocol: "invocations", Version: "v0.0.1"}, + }, + }, + { + name: "unknown protocol", + flagProtocols: []string{"unknown_proto"}, + wantErr: true, + wantErrContain: "unknown protocol", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got, err := promptProtocols(t.Context(), nil, false, tt.flagProtocols) + if tt.wantErr { + if err == nil { + t.Fatal("expected error, got nil") + } + if tt.wantErrContain != "" && !strings.Contains(err.Error(), tt.wantErrContain) { + t.Errorf("error = %q, want containing %q", err.Error(), tt.wantErrContain) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(got) != len(tt.wantProtocols) { + t.Fatalf("got %d protocols, want %d", len(got), len(tt.wantProtocols)) + } + for i := range got { + if got[i].Protocol != tt.wantProtocols[i].Protocol { + t.Errorf("protocol[%d] = %q, want %q", i, got[i].Protocol, tt.wantProtocols[i].Protocol) + } + if got[i].Version != tt.wantProtocols[i].Version { + t.Errorf("version[%d] = %q, want %q", i, got[i].Version, tt.wantProtocols[i].Version) + } + } + }) + } +} + +func TestPromptProtocols_NoPromptDefault(t *testing.T) { + t.Parallel() + + got, err := promptProtocols(t.Context(), nil, true, nil) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(got) != 1 { + t.Fatalf("got %d protocols, want 1", len(got)) + } + if got[0].Protocol != "responses" { + t.Errorf("protocol = %q, want %q", got[0].Protocol, "responses") + } + if got[0].Version != "v1" { + t.Errorf("version = %q, want %q", got[0].Version, "v1") + } +} + +func TestKnownProtocolNames(t *testing.T) { + t.Parallel() + + result := knownProtocolNames() + if !strings.Contains(result, "responses") { + t.Errorf("knownProtocolNames() = %q, want to contain 'responses'", result) + } + if !strings.Contains(result, "invocations") { + t.Errorf("knownProtocolNames() = %q, want to contain 'invocations'", result) + } +} From c239f27b66ffb3036bf2059a95ad49ec42f48a63 Mon Sep 17 00:00:00 2001 From: trangevi Date: Fri, 3 Apr 2026 11:01:53 -0700 Subject: [PATCH 2/2] PR comments Signed-off-by: trangevi --- .../internal/cmd/init_from_code.go | 29 ++-- .../internal/cmd/init_from_code_test.go | 132 ++++++++++++++++++ 2 files changed, 153 insertions(+), 8 deletions(-) diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go index 33b18ab69cc..bba299b8c5c 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code.go @@ -443,7 +443,7 @@ func (a *InitFromCodeAction) createDefinitionFromLocalAgent(ctx context.Context) agentKind := agent_yaml.AgentKindHosted // Prompt user for supported protocols - protocols, err := promptProtocols(ctx, a.azdClient, a.flags.NoPrompt, a.flags.protocols) + protocols, err := promptProtocols(ctx, a.azdClient.Prompt(), a.flags.NoPrompt, a.flags.protocols) if err != nil { return nil, err } @@ -806,7 +806,7 @@ var knownProtocols = []protocolInfo{ // When noPrompt is true and no flag values are provided, defaults to [responses/v1]. func promptProtocols( ctx context.Context, - azdClient *azdext.AzdClient, + promptClient azdext.PromptServiceClient, noPrompt bool, flagProtocols []string, ) ([]agent_yaml.ProtocolVersionRecord, error) { @@ -816,10 +816,16 @@ func promptProtocols( versionOf[p.Name] = p.Version } - // If explicit flag values were provided, use them directly. + // If explicit flag values were provided, use them directly (with dedup). if len(flagProtocols) > 0 { + seen := make(map[string]bool, len(flagProtocols)) records := make([]agent_yaml.ProtocolVersionRecord, 0, len(flagProtocols)) for _, name := range flagProtocols { + if seen[name] { + continue + } + seen[name] = true + version, ok := versionOf[name] if !ok { return nil, exterrors.Validation( @@ -854,11 +860,11 @@ func promptProtocols( }) } - resp, err := azdClient.Prompt().MultiSelect(ctx, &azdext.MultiSelectRequest{ + resp, err := promptClient.MultiSelect(ctx, &azdext.MultiSelectRequest{ Options: &azdext.MultiSelectOptions{ - Message: "Which protocols does your agent support?", - Choices: choices, - Hint: "Use arrow keys to move, space to toggle, enter to confirm", + Message: "Which protocols does your agent support?", + Choices: choices, + HelpMessage: "Use arrow keys to move, space to toggle, enter to confirm", }, }) if err != nil { @@ -872,9 +878,16 @@ func promptProtocols( var records []agent_yaml.ProtocolVersionRecord for _, choice := range resp.Values { if choice.Selected { + version, ok := versionOf[choice.Value] + if !ok { + return nil, exterrors.Internal( + "prompt_protocols", + fmt.Sprintf("unexpected protocol %q returned from prompt", choice.Value), + ) + } records = append(records, agent_yaml.ProtocolVersionRecord{ Protocol: choice.Value, - Version: versionOf[choice.Value], + Version: version, }) } } diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go index 6af3eb50aa3..968eabdefa7 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_from_code_test.go @@ -5,10 +5,16 @@ package cmd import ( "azureaiagent/internal/pkg/agents/agent_yaml" + "context" "os" "path/filepath" "strings" "testing" + + "github.com/azure/azure-dev/cli/azd/pkg/azdext" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" ) func TestSanitizeAgentName(t *testing.T) { @@ -619,6 +625,14 @@ func TestPromptProtocols_FlagValues(t *testing.T) { wantErr: true, wantErrContain: "unknown protocol", }, + { + name: "duplicates are removed", + flagProtocols: []string{"responses", "responses", "invocations"}, + wantProtocols: []agent_yaml.ProtocolVersionRecord{ + {Protocol: "responses", Version: "v1"}, + {Protocol: "invocations", Version: "v0.0.1"}, + }, + }, } for _, tt := range tests { @@ -682,3 +696,121 @@ func TestKnownProtocolNames(t *testing.T) { t.Errorf("knownProtocolNames() = %q, want to contain 'invocations'", result) } } + +// fakePromptClient is a lightweight test double for azdext.PromptServiceClient. +type fakePromptClient struct { + azdext.PromptServiceClient + multiSelectFn func( + ctx context.Context, + in *azdext.MultiSelectRequest, + opts ...grpc.CallOption, + ) (*azdext.MultiSelectResponse, error) +} + +func (f *fakePromptClient) MultiSelect( + ctx context.Context, + in *azdext.MultiSelectRequest, + opts ...grpc.CallOption, +) (*azdext.MultiSelectResponse, error) { + return f.multiSelectFn(ctx, in, opts...) +} + +func TestPromptProtocols_Interactive(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + multiSelectFn func(context.Context, *azdext.MultiSelectRequest, ...grpc.CallOption) (*azdext.MultiSelectResponse, error) + wantProtocols []agent_yaml.ProtocolVersionRecord + wantErr bool + wantErrContain string + }{ + { + name: "both protocols selected", + multiSelectFn: func(_ context.Context, _ *azdext.MultiSelectRequest, _ ...grpc.CallOption) (*azdext.MultiSelectResponse, error) { + return &azdext.MultiSelectResponse{ + Values: []*azdext.MultiSelectChoice{ + {Value: "responses", Label: "responses", Selected: true}, + {Value: "invocations", Label: "invocations", Selected: true}, + }, + }, nil + }, + wantProtocols: []agent_yaml.ProtocolVersionRecord{ + {Protocol: "responses", Version: "v1"}, + {Protocol: "invocations", Version: "v0.0.1"}, + }, + }, + { + name: "single protocol selected", + multiSelectFn: func(_ context.Context, _ *azdext.MultiSelectRequest, _ ...grpc.CallOption) (*azdext.MultiSelectResponse, error) { + return &azdext.MultiSelectResponse{ + Values: []*azdext.MultiSelectChoice{ + {Value: "responses", Label: "responses", Selected: true}, + {Value: "invocations", Label: "invocations", Selected: false}, + }, + }, nil + }, + wantProtocols: []agent_yaml.ProtocolVersionRecord{ + {Protocol: "responses", Version: "v1"}, + }, + }, + { + name: "user cancellation", + multiSelectFn: func(_ context.Context, _ *azdext.MultiSelectRequest, _ ...grpc.CallOption) (*azdext.MultiSelectResponse, error) { + return nil, status.Error(codes.Canceled, "cancelled by user") + }, + wantErr: true, + wantErrContain: "cancelled", + }, + { + name: "empty selection returns validation error", + multiSelectFn: func(_ context.Context, _ *azdext.MultiSelectRequest, _ ...grpc.CallOption) (*azdext.MultiSelectResponse, error) { + return &azdext.MultiSelectResponse{ + Values: []*azdext.MultiSelectChoice{ + {Value: "responses", Label: "responses", Selected: false}, + {Value: "invocations", Label: "invocations", Selected: false}, + }, + }, nil + }, + wantErr: true, + wantErrContain: "at least one protocol must be selected", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + client := &fakePromptClient{multiSelectFn: tt.multiSelectFn} + got, err := promptProtocols(t.Context(), client, false, nil) + if tt.wantErr { + if err == nil { + t.Fatal("expected error, got nil") + } + if tt.wantErrContain != "" && + !strings.Contains(err.Error(), tt.wantErrContain) { + t.Errorf("error = %q, want containing %q", + err.Error(), tt.wantErrContain) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(got) != len(tt.wantProtocols) { + t.Fatalf("got %d protocols, want %d", + len(got), len(tt.wantProtocols)) + } + for i := range got { + if got[i].Protocol != tt.wantProtocols[i].Protocol { + t.Errorf("protocol[%d] = %q, want %q", + i, got[i].Protocol, tt.wantProtocols[i].Protocol) + } + if got[i].Version != tt.wantProtocols[i].Version { + t.Errorf("version[%d] = %q, want %q", + i, got[i].Version, tt.wantProtocols[i].Version) + } + } + }) + } +}