diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/hosted-agent-regions.json b/cli/azd/extensions/azure.ai.agents/internal/cmd/hosted-agent-regions.json new file mode 100644 index 00000000000..c5b98a6b034 --- /dev/null +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/hosted-agent-regions.json @@ -0,0 +1,22 @@ +{ + "regions": [ + "australiaeast", + "brazilsouth", + "canadacentral", + "eastus2", + "francecentral", + "japaneast", + "koreacentral", + "northcentralus", + "norwayeast", + "polandcentral", + "southafricanorth", + "southeastasia", + "southindia", + "spaincentral", + "swedencentral", + "switzerlandnorth", + "westus", + "westus3" + ] +} diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_foundry_resources_helpers.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_foundry_resources_helpers.go index b8f94476a6b..285de31c6d4 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_foundry_resources_helpers.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_foundry_resources_helpers.go @@ -735,7 +735,10 @@ func ensureLocation( azureContext *azdext.AzureContext, envName string, ) error { - allowedLocations := supportedRegionsForInit() + allowedLocations, err := supportedRegionsForInit(ctx) + if err != nil { + return err + } if azureContext.Scope.Location != "" && locationAllowed(azureContext.Scope.Location, allowedLocations) { return nil diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations.go index a796f202e59..b8044b87208 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations.go @@ -3,49 +3,225 @@ package cmd -import "slices" - -// No API available to query supported regions for hosted agents, so keep hardcoded list based on public documentation: -// https://learn.microsoft.com/azure/foundry/agents/concepts/hosted-agents#region-availability -var supportedHostedAgentRegions = []string{ - "australiaeast", - "brazilsouth", - "canadacentral", - "canadaeast", - "centralus", - "eastus", - "eastus2", - "francecentral", - "germanywestcentral", - "italynorth", - "japaneast", - "koreacentral", - "northcentralus", - "norwayeast", - "polandcentral", - "southafricanorth", - "southcentralus", - "southeastasia", - "southindia", - "spaincentral", - "swedencentral", - "switzerlandnorth", - "uaenorth", - "uksouth", - "westeurope", - "westus", - "westus3", -} - -func supportedRegionsForInit() []string { - return slices.Clone(supportedHostedAgentRegions) -} - -// supportedModelLocations returns the intersection of a model's available locations -// with the supported hosted agent regions. -func supportedModelLocations(modelLocations []string) []string { - supported := supportedRegionsForInit() - return slices.DeleteFunc(slices.Clone(modelLocations), func(loc string) bool { +import ( + "context" + _ "embed" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "slices" + "sync" + "time" + + "azureaiagent/internal/exterrors" + + "github.com/azure/azure-dev/cli/azd/pkg/azdext" +) + +// hostedAgentRegionsURL points at the supported-regions manifest. +// It is a var so tests can override it. +var hostedAgentRegionsURL = "https://aka.ms/azd-ai-agents/regions" + +// embeddedHostedAgentRegionsJSON is the build-time fallback used when the live +// manifest fetch fails (e.g. transient network issues, restrictive proxies). +// +//go:embed hosted-agent-regions.json +var embeddedHostedAgentRegionsJSON []byte + +const ( + hostedAgentRegionsFetchTimeout = 5 * time.Second + // hostedAgentRegionsManifestMaxBytes caps the manifest body to guard against + // unexpectedly large responses from the source URL. + hostedAgentRegionsManifestMaxBytes = 1 << 20 // 1 MiB +) + +type hostedAgentRegionsManifest struct { + Regions []string `json:"regions"` +} + +var regionsCache struct { + mu sync.Mutex + regions []string + inflight *regionsFetch +} + +// regionsFetch coordinates concurrent callers waiting on the same in-flight fetch +// so the package-level mutex can be released while the network call is running. +type regionsFetch struct { + done chan struct{} + regions []string + err error +} + +// supportedRegionsForInit returns the list of Azure regions supported for hosted agents. +// The result is cached for the process after the first successful fetch. +// +// The fetch itself is performed without holding regionsCache.mu so callers whose +// context is canceled can return promptly even if another goroutine is mid-fetch. +func supportedRegionsForInit(ctx context.Context) ([]string, error) { + regionsCache.mu.Lock() + if regionsCache.regions != nil { + regions := slices.Clone(regionsCache.regions) + regionsCache.mu.Unlock() + return regions, nil + } + + fetch := regionsCache.inflight + if fetch == nil { + fetch = ®ionsFetch{done: make(chan struct{})} + regionsCache.inflight = fetch + // context.WithoutCancel keeps any context values but drops cancellation, + // because the fetch result is shared across all waiters and must not be + // aborted by a single caller's cancellation. + go runRegionsFetch(context.WithoutCancel(ctx), fetch) + } + regionsCache.mu.Unlock() + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-fetch.done: + if fetch.err != nil { + return nil, fetch.err + } + return slices.Clone(fetch.regions), nil + } +} + +// runRegionsFetch performs the network fetch, populates the cache on success, and +// signals all waiters via fetch.done. If the fetch fails, the embedded build-time +// manifest is used as a fallback so a transient network issue doesn't halt init. +// +// ctx must not carry a cancellation that any single caller can trigger, since the +// fetch result is shared. Callers pass context.WithoutCancel(callerCtx). +func runRegionsFetch(ctx context.Context, fetch *regionsFetch) { + // The fetch applies its own timeout (hostedAgentRegionsFetchTimeout). + regions, err := fetchHostedAgentRegionsFromURL(ctx, http.DefaultClient, hostedAgentRegionsURL) + + if err != nil { + if fallback, fbErr := parseEmbeddedHostedAgentRegions(); fbErr == nil && len(fallback) > 0 { + regions = fallback + err = nil + } + } + + regionsCache.mu.Lock() + if err == nil { + regionsCache.regions = regions + } + regionsCache.inflight = nil + regionsCache.mu.Unlock() + + fetch.regions = regions + fetch.err = err + close(fetch.done) +} + +// parseEmbeddedHostedAgentRegions decodes the embedded build-time manifest used +// as a fallback when the live fetch fails. +func parseEmbeddedHostedAgentRegions() ([]string, error) { + var manifest hostedAgentRegionsManifest + if err := json.Unmarshal(embeddedHostedAgentRegionsJSON, &manifest); err != nil { + return nil, err + } + regions := make([]string, 0, len(manifest.Regions)) + for _, r := range manifest.Regions { + if normalized := normalizeLocationName(r); normalized != "" { + regions = append(regions, normalized) + } + } + return regions, nil +} + +// supportedModelLocations returns the intersection of a model's available locations with +// the supported hosted-agent regions. Returns an error when the intersection is empty +// because passing an empty allowlist downstream disables filtering, which would let users +// pick regions that are not supported for hosted agents. +func supportedModelLocations(ctx context.Context, modelLocations []string) ([]string, error) { + supported, err := supportedRegionsForInit(ctx) + if err != nil { + return nil, err + } + + result := slices.DeleteFunc(slices.Clone(modelLocations), func(loc string) bool { return !locationAllowed(loc, supported) }) + + if len(result) == 0 { + return nil, exterrors.Dependency( + exterrors.CodeNoSupportedModelLocations, + "the selected model is not available in any region supported for hosted agents", + "select a different model.", + ) + } + + return result, nil +} + +func fetchHostedAgentRegionsFromURL(ctx context.Context, httpClient *http.Client, url string) ([]string, error) { + fetchCtx, cancel := context.WithTimeout(ctx, hostedAgentRegionsFetchTimeout) + defer cancel() + + req, err := http.NewRequestWithContext(fetchCtx, http.MethodGet, url, nil) + if err != nil { + return nil, regionsFetchError(err) + } + + //nolint:gosec // URL is the hardcoded hostedAgentRegionsURL constant or test override + resp, err := httpClient.Do(req) + if err != nil { + return nil, regionsFetchError(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, regionsFetchError(fmt.Errorf("unexpected HTTP status %d", resp.StatusCode)) + } + + body, err := io.ReadAll(io.LimitReader(resp.Body, hostedAgentRegionsManifestMaxBytes+1)) + if err != nil { + return nil, regionsFetchError(err) + } + if len(body) > hostedAgentRegionsManifestMaxBytes { + return nil, regionsFetchError(fmt.Errorf( + "manifest exceeds %d byte limit", hostedAgentRegionsManifestMaxBytes, + )) + } + + var manifest hostedAgentRegionsManifest + if err := json.Unmarshal(body, &manifest); err != nil { + return nil, regionsFetchError(err) + } + + regions := make([]string, 0, len(manifest.Regions)) + for _, r := range manifest.Regions { + if normalized := normalizeLocationName(r); normalized != "" { + regions = append(regions, normalized) + } + } + + if len(regions) == 0 { + return nil, regionsFetchError(fmt.Errorf("manifest contained no valid regions")) + } + + return regions, nil +} + +func regionsFetchError(err error) error { + return exterrors.Dependency( + exterrors.CodeRegionsFetchFailed, + fmt.Sprintf("could not retrieve the list of supported Azure regions: %v", err), + "check your network connection and try again. "+ + "If the issue persists, file an issue at https://github.com/Azure/azure-dev/issues", + ) +} + +// isNoSupportedLocationsError reports whether err is the structured error returned by +// [supportedModelLocations] when no region in the model's location list is supported +// for hosted agents. +func isNoSupportedLocationsError(err error) bool { + localErr, ok := errors.AsType[*azdext.LocalError](err) + return ok && localErr.Code == exterrors.CodeNoSupportedModelLocations } diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations_test.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations_test.go index 8e0bd8bab71..8562a1d3bc5 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations_test.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_locations_test.go @@ -4,78 +4,261 @@ package cmd import ( + "errors" + "net/http" + "net/http/httptest" + "slices" + "strings" + "sync" + "sync/atomic" "testing" + "time" + "github.com/azure/azure-dev/cli/azd/pkg/azdext" "github.com/stretchr/testify/require" + + "azureaiagent/internal/exterrors" ) -func TestSupportedModelLocations(t *testing.T) { +func TestFetchHostedAgentRegionsFromURL_Success(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"regions": ["eastus2", "westus3", "swedencentral"]}`)) + })) + t.Cleanup(server.Close) + + regions, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + require.NoError(t, err) + require.Equal(t, []string{"eastus2", "westus3", "swedencentral"}, regions) +} + +func TestFetchHostedAgentRegionsFromURL_NormalizesEntries(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"regions": [" EastUS2 ", "westus3", "", " "]}`)) + })) + t.Cleanup(server.Close) + + regions, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + require.NoError(t, err) + require.Equal(t, []string{"eastus2", "westus3"}, regions) +} + +func TestFetchHostedAgentRegionsFromURL_HTTPError(t *testing.T) { t.Parallel() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "boom", http.StatusInternalServerError) + })) + t.Cleanup(server.Close) + + _, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + require.Error(t, err) +} + +func TestFetchHostedAgentRegionsFromURL_MalformedJSON(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{not json`)) + })) + t.Cleanup(server.Close) + + _, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + require.Error(t, err) +} + +func TestFetchHostedAgentRegionsFromURL_EmptyManifest(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"regions": []}`)) + })) + t.Cleanup(server.Close) + + _, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + require.Error(t, err) +} + +func TestFetchHostedAgentRegionsFromURL_RespectsTimeout(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + time.Sleep(hostedAgentRegionsFetchTimeout + 2*time.Second) + })) + t.Cleanup(server.Close) + + start := time.Now() + _, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + elapsed := time.Since(start) + + require.Error(t, err) + require.Less(t, elapsed, hostedAgentRegionsFetchTimeout+1*time.Second) +} + +func TestSupportedModelLocations(t *testing.T) { + resetRegionsCache(t, []string{"eastus2", "westus3"}) + tests := []struct { name string modelLocations []string - wantSubset bool - wantLen int + want []string + wantErr bool }{ - { - name: "AllSupported", - modelLocations: []string{"eastus", "westus"}, - wantSubset: true, - wantLen: 2, - }, - { - name: "SomeUnsupported", - modelLocations: []string{"eastus", "unsupportedregion"}, - wantSubset: true, - wantLen: 1, - }, - { - name: "NoneSupported", - modelLocations: []string{"unsupportedregion1", "unsupportedregion2"}, - wantSubset: true, - wantLen: 0, - }, - { - name: "EmptyInput", - modelLocations: []string{}, - wantSubset: true, - wantLen: 0, - }, - { - name: "NilInput", - modelLocations: nil, - wantSubset: true, - wantLen: 0, - }, + {"AllSupported", []string{"eastus2", "westus3"}, []string{"eastus2", "westus3"}, false}, + {"SomeUnsupported", []string{"eastus2", "unsupported"}, []string{"eastus2"}, false}, + {"NoneSupported", []string{"unsupported1", "unsupported2"}, nil, true}, + {"EmptyInput", []string{}, nil, true}, + {"NilInput", nil, nil, true}, } - supported := supportedRegionsForInit() - for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - result := supportedModelLocations(tt.modelLocations) - require.Len(t, result, tt.wantLen) - - // Every returned location must be in the supported regions list - for _, loc := range result { - require.True(t, locationAllowed(loc, supported), - "returned location %q should be in supported regions", loc) + result, err := supportedModelLocations(t.Context(), tt.modelLocations) + if tt.wantErr { + require.Error(t, err) + return } + require.NoError(t, err) + require.ElementsMatch(t, tt.want, result) }) } } -func TestSupportedModelLocationsDoesNotMutateInput(t *testing.T) { +func TestSupportedModelLocations_EmptyIntersectionReturnsStructuredError(t *testing.T) { + resetRegionsCache(t, []string{"eastus2"}) + + _, err := supportedModelLocations(t.Context(), []string{"unsupported"}) + require.Error(t, err) + localErr, ok := errors.AsType[*azdext.LocalError](err) + require.True(t, ok, "expected *azdext.LocalError, got %T", err) + require.Equal(t, exterrors.CodeNoSupportedModelLocations, localErr.Code) +} + +func TestSupportedModelLocations_DoesNotMutateInput(t *testing.T) { + resetRegionsCache(t, []string{"eastus2", "westus3"}) + + input := []string{"eastus2", "unsupported", "westus3"} + original := slices.Clone(input) + + _, err := supportedModelLocations(t.Context(), input) + require.NoError(t, err) + require.Equal(t, original, input) +} + +func TestSupportedRegionsForInit_FetchesOnceAndCaches(t *testing.T) { + resetRegionsCache(t, nil) + + hits := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits++ + _, _ = w.Write([]byte(`{"regions": ["eastus2"]}`)) + })) + t.Cleanup(server.Close) + + prev := hostedAgentRegionsURL + hostedAgentRegionsURL = server.URL + t.Cleanup(func() { hostedAgentRegionsURL = prev }) + + for range 3 { + got, err := supportedRegionsForInit(t.Context()) + require.NoError(t, err) + require.Equal(t, []string{"eastus2"}, got) + } + require.Equal(t, 1, hits) +} + +func TestFetchHostedAgentRegionsFromURL_RejectsOversizedManifest(t *testing.T) { + t.Parallel() + + huge := strings.Repeat("a", hostedAgentRegionsManifestMaxBytes+1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"regions":["` + huge + `"]}`)) + })) + t.Cleanup(server.Close) + + _, err := fetchHostedAgentRegionsFromURL(t.Context(), http.DefaultClient, server.URL) + require.Error(t, err) + require.Contains(t, err.Error(), "byte limit") +} + +func TestSupportedRegionsForInit_ConcurrentCallersFetchOnce(t *testing.T) { + resetRegionsCache(t, nil) + + var hits atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + // Brief delay so concurrent callers genuinely overlap on the in-flight fetch. + time.Sleep(50 * time.Millisecond) + _, _ = w.Write([]byte(`{"regions": ["eastus2"]}`)) + })) + t.Cleanup(server.Close) + + prev := hostedAgentRegionsURL + hostedAgentRegionsURL = server.URL + t.Cleanup(func() { hostedAgentRegionsURL = prev }) + + const callers = 8 + var wg sync.WaitGroup + wg.Add(callers) + for range callers { + go func() { + defer wg.Done() + got, err := supportedRegionsForInit(t.Context()) + require.NoError(t, err) + require.Equal(t, []string{"eastus2"}, got) + }() + } + wg.Wait() + + require.Equal(t, int32(1), hits.Load()) +} + +func TestSupportedRegionsForInit_FallsBackToEmbeddedOnFetchError(t *testing.T) { + resetRegionsCache(t, nil) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "boom", http.StatusInternalServerError) + })) + t.Cleanup(server.Close) + + prev := hostedAgentRegionsURL + hostedAgentRegionsURL = server.URL + t.Cleanup(func() { hostedAgentRegionsURL = prev }) + + got, err := supportedRegionsForInit(t.Context()) + require.NoError(t, err) + + want, err := parseEmbeddedHostedAgentRegions() + require.NoError(t, err) + require.NotEmpty(t, want) + require.Equal(t, want, got) +} + +func TestParseEmbeddedHostedAgentRegions_NotEmpty(t *testing.T) { t.Parallel() - input := []string{"eastus", "unsupportedregion", "westus"} - original := make([]string, len(input)) - copy(original, input) + regions, err := parseEmbeddedHostedAgentRegions() + require.NoError(t, err) + require.NotEmpty(t, regions, "embedded fallback manifest must contain at least one region") +} + +func resetRegionsCache(t *testing.T, regions []string) { + t.Helper() - _ = supportedModelLocations(input) + regionsCache.mu.Lock() + prev := regionsCache.regions + prevInflight := regionsCache.inflight + regionsCache.regions = regions + regionsCache.inflight = nil + regionsCache.mu.Unlock() - require.Equal(t, original, input, "input slice should not be mutated") + t.Cleanup(func() { + regionsCache.mu.Lock() + regionsCache.regions = prev + regionsCache.inflight = prevInflight + regionsCache.mu.Unlock() + }) } diff --git a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_models.go b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_models.go index d31bc6aa410..1064359787b 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/cmd/init_models.go +++ b/cli/azd/extensions/azure.ai.agents/internal/cmd/init_models.go @@ -621,11 +621,22 @@ func (a *modelSelector) promptForModelLocationMismatch( } if selectedChoice == "location" { + allowedLocations, err := supportedModelLocations(ctx, currentModel.Locations) + if err != nil { + if isNoSupportedLocationsError(err) { + message = fmt.Sprintf( + "Model '%s' is not available in any region supported for hosted agents.", + currentModel.Name, + ) + continue + } + return nil, "", err + } locationResp, err := a.azdClient.Prompt().PromptAiModelLocationWithQuota(ctx, &azdext.PromptAiModelLocationWithQuotaRequest{ AzureContext: a.azureContext, ModelName: currentModel.Name, - AllowedLocations: supportedModelLocations(currentModel.Locations), + AllowedLocations: allowedLocations, Quota: &azdext.QuotaCheckOptions{ MinRemainingCapacity: 1, }, @@ -672,11 +683,23 @@ func (a *modelSelector) promptForModelLocationMismatch( } selectedModel := modelResp.Model + allowedLocations, err := supportedModelLocations(ctx, selectedModel.Locations) + if err != nil { + if isNoSupportedLocationsError(err) { + currentModel = selectedModel + message = fmt.Sprintf( + "Model '%s' is not available in any region supported for hosted agents.", + selectedModel.Name, + ) + continue + } + return nil, "", err + } locationResp, err := a.azdClient.Prompt().PromptAiModelLocationWithQuota(ctx, &azdext.PromptAiModelLocationWithQuotaRequest{ AzureContext: a.azureContext, ModelName: selectedModel.Name, - AllowedLocations: supportedModelLocations(selectedModel.Locations), + AllowedLocations: allowedLocations, Quota: &azdext.QuotaCheckOptions{ MinRemainingCapacity: 1, }, diff --git a/cli/azd/extensions/azure.ai.agents/internal/exterrors/codes.go b/cli/azd/extensions/azure.ai.agents/internal/exterrors/codes.go index 0fe0d309bdd..3a1714c6929 100644 --- a/cli/azd/extensions/azure.ai.agents/internal/exterrors/codes.go +++ b/cli/azd/extensions/azure.ai.agents/internal/exterrors/codes.go @@ -82,8 +82,10 @@ const ( // Used as fallback codes with [FromAiService] when the gRPC response // doesn't include a more specific ErrorInfo reason. const ( - CodeModelCatalogFailed = "model_catalog_failed" - CodeModelResolutionFailed = "model_resolution_failed" + CodeModelCatalogFailed = "model_catalog_failed" + CodeModelResolutionFailed = "model_resolution_failed" + CodeRegionsFetchFailed = "regions_fetch_failed" + CodeNoSupportedModelLocations = "no_supported_model_locations" ) // Error codes for session errors.