diff --git a/cli/azd/extensions/azure.ai.rle/CHANGELOG.md b/cli/azd/extensions/azure.ai.rle/CHANGELOG.md index d99903acb97..f50a9b523a9 100644 --- a/cli/azd/extensions/azure.ai.rle/CHANGELOG.md +++ b/cli/azd/extensions/azure.ai.rle/CHANGELOG.md @@ -1,6 +1,6 @@ # Release History -## 0.4.0-preview (Unreleased) +## 0.4.1-preview (Unreleased) - Align environment discovery and remote invocation with the refreshed RLE service routes and cursor-based response contracts. - Use `/rl_environments` consistently for environment and instance lifecycle APIs. @@ -10,6 +10,7 @@ - Delete the temporary instance and group on exit with Ctrl+C-independent cleanup and concise terminal status. - Persist the environment name as `environmentName` while continuing to read legacy `name` state files. - Authenticate and API-version OpenEnv gateway requests on the configured Foundry project origin, wait for runtime health before reporting readiness, and route the browser playground through an authenticated local proxy. +- Initialize a required local folder by interactively selecting and sparsely downloading an environment from the RLE samples repository. ## 0.3.0-preview diff --git a/cli/azd/extensions/azure.ai.rle/README.md b/cli/azd/extensions/azure.ai.rle/README.md index c614097bd85..05faf29ecde 100644 --- a/cli/azd/extensions/azure.ai.rle/README.md +++ b/cli/azd/extensions/azure.ai.rle/README.md @@ -10,7 +10,7 @@ Install: - Azure CLI (`az`): https://learn.microsoft.com/cli/azure/install-azure-cli - Docker Desktop: https://www.docker.com/products/docker-desktop/ - Go, if building from source: https://go.dev/doc/install -- Git, if building from source: https://git-scm.com/downloads +- Git, required by `azd ai rle init` to download samples: https://git-scm.com/downloads Verify: @@ -81,23 +81,31 @@ az acr login --name ### 1. Initialize an environment session -Default echo session: +Select a sample and use its name for the local folder and RLE environment: ```powershell azd ai rle init -cd .\echo_env ``` -The default echo session downloads the Hugging Face `OpenEnv` repo, copies `envs/echo_env` into the session folder, and writes `.azd-rle.json` with the local `environmentName`. Existing state files that use the legacy `name` property remain supported. +`init` reads the available environments from +[rle-samples](https://github.com/sujit-kamireddy/rle-samples) and prompts you to select one. +Only the selected sample is downloaded. For example, selecting `code_rl` copies it into `.\code_rl` +and stores `code_rl` as the RLE environment name in `.azd-rle.json`. -The copied session does not keep `.git` metadata from the upstream repository. +To use a different local folder and RLE environment name, provide it before selecting a sample: -Name the copied echo session: +```powershell +azd ai rle init my_environment +``` + +When prompts are disabled, the required positional name selects the sample and is also used for the folder: ```powershell -azd ai rle init code_rl +azd ai rle init code_rl --no-prompt ``` +The copied session does not keep `.git` metadata from the sample repository. + For an existing source folder, skip `init` and run commands directly from that folder. ### 2. Run locally diff --git a/cli/azd/extensions/azure.ai.rle/extension.yaml b/cli/azd/extensions/azure.ai.rle/extension.yaml index 93c51754c91..c19bea8bde0 100644 --- a/cli/azd/extensions/azure.ai.rle/extension.yaml +++ b/cli/azd/extensions/azure.ai.rle/extension.yaml @@ -11,11 +11,11 @@ tags: - ai - rle usage: azd ai rle [options] -version: 0.4.0-preview +version: 0.4.1-preview examples: - name: init - description: Copy the OpenEnv echo sample into a local RLE environment. - usage: azd ai rle init + description: Select an RLE sample and copy it into a new local folder. + usage: azd ai rle init [folder-name] - name: publish description: Build, push, and create or update the RLE environment. usage: azd ai rle publish --version-bump major diff --git a/cli/azd/extensions/azure.ai.rle/go.mod b/cli/azd/extensions/azure.ai.rle/go.mod index efe393449f5..31998eee0a7 100644 --- a/cli/azd/extensions/azure.ai.rle/go.mod +++ b/cli/azd/extensions/azure.ai.rle/go.mod @@ -7,6 +7,7 @@ require ( github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 github.com/azure/azure-dev/cli/azd v1.25.0 github.com/fatih/color v1.18.0 + github.com/gorilla/websocket v1.5.3 github.com/spf13/cobra v1.10.1 ) diff --git a/cli/azd/extensions/azure.ai.rle/go.sum b/cli/azd/extensions/azure.ai.rle/go.sum index e2fb881cb18..9dcf6549391 100644 --- a/cli/azd/extensions/azure.ai.rle/go.sum +++ b/cli/azd/extensions/azure.ai.rle/go.sum @@ -121,6 +121,8 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= github.com/hexops/gotextdiff v1.0.3/go.mod h1:pSWU5MAI3yDq+fZBTazCSJysOMbxWL1BSow5/V2vxeg= github.com/hinshun/vt10x v0.0.0-20220119200601-820417d04eec h1:qv2VnGeEQHchGaZ/u7lxST/RaJw+cv273q79D81Xbog= diff --git a/cli/azd/extensions/azure.ai.rle/internal/cmd/init.go b/cli/azd/extensions/azure.ai.rle/internal/cmd/init.go index 901d92c4d1a..61719371e9b 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/cmd/init.go +++ b/cli/azd/extensions/azure.ai.rle/internal/cmd/init.go @@ -4,9 +4,11 @@ package cmd import ( + "context" "fmt" "os" "runtime" + "slices" "strings" "azure.ai.rle/internal/project" @@ -20,34 +22,50 @@ type rleInitFlags struct { } type initAction struct { - cmd *cobra.Command - flags *rleInitFlags - envNameOverride string + cmd *cobra.Command + flags *rleInitFlags + folderName string + noPrompt bool } -var checkoutOpenEnvEchoSampleFunc = project.CheckoutOpenEnvEchoSample +type rleSampleCatalog interface { + SampleNames() []string + Copy(sampleName string, folderName string, dest string, force bool) (string, error) + Close() error +} + +var loadRleSampleCatalogFunc = func() (rleSampleCatalog, error) { + return project.LoadRleSampleCatalog() +} -func newInitCommand() *cobra.Command { +var selectRleSampleFunc = selectRleSample + +func newInitCommand(noPrompt *bool) *cobra.Command { flags := &rleInitFlags{} cmd := &cobra.Command{ - Use: "init", - Short: "Initialize a local RLE environment", + Use: "init [folder-name]", + Short: "Initialize a local RLE environment from a sample", Args: cobra.MaximumNArgs(1), RunE: func(cmd *cobra.Command, args []string) error { - envNameOverride := "" + folderName := "" if len(args) == 1 { - envNameOverride = args[0] + folderName = args[0] } - return (&initAction{cmd: cmd, flags: flags, envNameOverride: envNameOverride}).Run() + return (&initAction{ + cmd: cmd, + flags: flags, + folderName: folderName, + noPrompt: noPrompt != nil && *noPrompt, + }).Run() }, } cmd.SetHelpFunc(func(cmd *cobra.Command, args []string) { var help strings.Builder - help.WriteString("Initialize a local RLE environment\n") + help.WriteString("Initialize a local RLE environment from a sample\n") help.WriteString("Usage:\n") - help.WriteString(" rle init [environment-name] [flags]\n") + help.WriteString(" rle init [folder-name] [flags]\n") help.WriteString("Flags:\n") help.WriteString(" --force Overwrite generated files in an existing non-empty session directory\n") help.WriteString(" -h, --help help for init\n") @@ -62,32 +80,125 @@ func newInitCommand() *cobra.Command { } func (a *initAction) Run() error { - envName := firstNonEmpty(a.envNameOverride, "echo_env") - var err error - envName, err = project.ValidateEnvironmentName(envName) - if err != nil { + if a.noPrompt && a.folderName == "" { return &azdext.LocalError{ - Message: err.Error(), - Code: "rle_invalid_environment_name", + Message: "A sample name is required when prompts are disabled.", + Code: "rle_sample_name_required", Category: azdext.LocalErrorCategoryUser, - Suggestion: "Use snake_case starting with a letter, for example code_rl.", + Suggestion: "Run azd ai rle init --no-prompt.", + } + } + folderName := a.folderName + if folderName != "" { + var err error + folderName, err = validateRleFolderName(folderName) + if err != nil { + return err } } - sessionDir, err := checkoutOpenEnvEchoSampleFunc(envName, ".", a.flags.force) + catalog, err := loadRleSampleCatalogFunc() + if err != nil { + return err + } + defer func() { + _ = catalog.Close() + }() + requestedSample := "" + if a.noPrompt { + requestedSample = folderName + } + sampleName, err := resolveRleSample(a.cmd.Context(), requestedSample, catalog.SampleNames()) + if err != nil { + return err + } + if folderName == "" { + folderName, err = validateRleFolderName(sampleName) + if err != nil { + return err + } + } + sessionDir, err := catalog.Copy(sampleName, folderName, ".", a.flags.force) if err != nil { return err } - if err := saveRleStateIn(sessionDir, defaultRleState(envName)); err != nil { + if err := saveRleStateIn(sessionDir, defaultRleState(folderName)); err != nil { return err } displayDir := "." + string(os.PathSeparator) + sessionDir + if _, err := fmt.Fprintf(a.cmd.OutOrStdout(), "Copied RLE sample %q.\n", sampleName); err != nil { + return err + } _, err = fmt.Fprint(a.cmd.OutOrStdout(), initNextSteps(displayDir, runtime.GOOS, os.Getenv("SHELL"))) return err } +func validateRleFolderName(folderName string) (string, error) { + validated, err := project.ValidateEnvironmentName(folderName) + if err == nil { + return validated, nil + } + return "", &azdext.LocalError{ + Message: err.Error(), + Code: "rle_invalid_environment_name", + Category: azdext.LocalErrorCategoryUser, + Suggestion: "Use snake_case starting with a letter, for example code_rl.", + } +} + +func resolveRleSample(ctx context.Context, requestedSample string, sampleNames []string) (string, error) { + if requestedSample == "" { + return selectRleSampleFunc(ctx, sampleNames) + } + if slices.Contains(sampleNames, requestedSample) { + return requestedSample, nil + } + return "", &azdext.LocalError{ + Message: fmt.Sprintf("RLE sample %q was not found.", requestedSample), + Code: "rle_sample_not_found", + Category: azdext.LocalErrorCategoryUser, + Suggestion: fmt.Sprintf("Choose one of the available samples: %s.", strings.Join(sampleNames, ", ")), + } +} + +func selectRleSample(ctx context.Context, sampleNames []string) (string, error) { + if len(sampleNames) == 0 { + return "", &azdext.LocalError{ + Message: "No RLE samples are available.", + Code: "rle_samples_empty", + Category: azdext.LocalErrorCategoryUser, + Suggestion: "Add a sample to the RLE samples repository, then retry.", + } + } + choices := make([]*azdext.SelectChoice, len(sampleNames)) + for index, sampleName := range sampleNames { + choices[index] = &azdext.SelectChoice{Label: sampleName, Value: sampleName} + } + azdClient, err := azdext.NewAzdClient() + if err != nil { + return "", fmt.Errorf("create azd client for sample selection: %w", err) + } + defer azdClient.Close() + response, err := azdClient.Prompt().Select(azdext.WithAccessToken(ctx), &azdext.SelectRequest{ + Options: &azdext.SelectOptions{ + Message: "Select an RLE sample", + Choices: choices, + DisplayNumbers: new(true), + EnableFiltering: new(true), + }, + }) + if err != nil { + return "", fmt.Errorf("select RLE sample: %w", err) + } + selectedIndex := int(response.GetValue()) + if selectedIndex < 0 || selectedIndex >= len(sampleNames) { + return "", fmt.Errorf("invalid RLE sample selection index: %d", selectedIndex) + } + return sampleNames[selectedIndex], nil +} + func initNextSteps(displayDir string, goos string, shell string) string { projectEndpoint := `https://.services.ai.azure.com/api/projects/` registryEndpoint := `.azurecr.io` @@ -111,7 +222,7 @@ func initNextSteps(displayDir string, goos string, shell string) string { } return fmt.Sprintf( - "Created OpenEnv-style environment at: %s\n"+ + "Created RLE environment at: %s\n"+ "\nRun locally:\n"+ " cd \"%s\"\n"+ " azd ai rle run\n"+ diff --git a/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke.go b/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke.go index 989e21c9f38..89778ad39fc 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke.go +++ b/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke.go @@ -7,6 +7,7 @@ import ( "context" "crypto/rand" "encoding/hex" + "encoding/json" "errors" "fmt" "io" @@ -40,7 +41,7 @@ var validateSandboxURL = validateRemoteSandboxURL func newInvokeCommand() *cobra.Command { flags := &remoteInvokeFlags{ - timeout: 30, + timeout: 60, } cmd := &cobra.Command{ @@ -147,10 +148,17 @@ func (a *remoteInvokeAction) Run() error { ); err != nil { return err } + runtimeSession := project.NewWebSocketRuntimeSession( + instanceUrl, + a.flags.timeout, + client.authorizationHeader, + ) + defer runtimeSession.Close() playgroundUrl, stopPlayground, err := remotePlaygroundUrlWithAuthorizationProvider( ctx, instanceUrl, client.authorizationHeader, + runtimeSession, ) if err != nil { return err @@ -159,13 +167,11 @@ func (a *remoteInvokeAction) Run() error { if err := ui.OpenBrowser(playgroundUrl); err != nil { _, _ = fmt.Fprintf(a.cmd.ErrOrStderr(), "Warning: failed to open playground UI: %v\n", err) } - return project.RunShellWithContextAndAuthorizationProvider( + return project.RunWebSocketShellWithSession( ctx, a.cmd.InOrStdin(), a.cmd.OutOrStdout(), - instanceUrl, - a.flags.timeout, - client.authorizationHeader, + runtimeSession, ) } @@ -500,6 +506,7 @@ func remotePlaygroundUrlWithAuthorizationProvider( ctx context.Context, sandboxUrl string, authorizationProvider project.AuthorizationProvider, + runtimeSessions ...*project.WebSocketRuntimeSession, ) (string, func(), error) { listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { @@ -517,6 +524,7 @@ func remotePlaygroundUrlWithAuthorizationProvider( authorizationProvider, listener.Addr().String(), sessionToken, + runtimeSessions..., ), ReadHeaderTimeout: 5 * time.Second, } @@ -544,6 +552,7 @@ func remotePlaygroundHandler( authorizationProvider project.AuthorizationProvider, expectedHost string, sessionToken string, + runtimeSessions ...*project.WebSocketRuntimeSession, ) http.Handler { mux := http.NewServeMux() mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { @@ -558,7 +567,7 @@ func remotePlaygroundHandler( _, _ = io.WriteString(w, ui.RemotePlaygroundHTML) return } - proxyOpenEnvToSandbox(w, r, sandboxUrl, authorizationProvider) + proxyOpenEnvToSandbox(w, r, sandboxUrl, authorizationProvider, runtimeSessions...) }) return mux } @@ -631,6 +640,7 @@ func proxyOpenEnvToSandbox( r *http.Request, sandboxUrl string, authorizationProvider project.AuthorizationProvider, + runtimeSessions ...*project.WebSocketRuntimeSession, ) { operation := strings.Trim(r.URL.Path, "/") switch operation { @@ -648,6 +658,11 @@ func proxyOpenEnvToSandbox( http.NotFound(w, r) return } + if len(runtimeSessions) > 0 && runtimeSessions[0] != nil && + (operation == "reset" || operation == "step" || operation == "state") { + proxyStatefulOpenEnvOperation(w, r, operation, runtimeSessions[0]) + return + } targetUrl, err := project.RuntimeOperationURL(sandboxUrl, operation) if err != nil { @@ -687,6 +702,40 @@ func proxyOpenEnvToSandbox( _, _ = io.Copy(w, resp.Body) } +func proxyStatefulOpenEnvOperation( + w http.ResponseWriter, + r *http.Request, + operation string, + runtimeSession *project.WebSocketRuntimeSession, +) { + payload := "" + if operation == "reset" || operation == "step" { + body, err := io.ReadAll(http.MaxBytesReader(w, r.Body, 100*1024*1024)) + if err != nil { + http.Error(w, "invalid request body", http.StatusBadRequest) + return + } + payload = string(body) + } + if operation == "step" { + var request struct { + Action json.RawMessage `json:"action"` + } + if err := json.Unmarshal([]byte(payload), &request); err != nil || len(request.Action) == 0 { + http.Error(w, "step requires an action", http.StatusBadRequest) + return + } + payload = string(request.Action) + } + response, err := runtimeSession.CallAndDrain(r.Context(), operation, payload) + if err != nil { + http.Error(w, err.Error(), http.StatusBadGateway) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, response) //nolint:gosec // The response is served as JSON, not executable HTML. +} + func withFoundryAPIVersion(runtimeUrl string) (string, error) { parsedUrl, err := url.Parse(runtimeUrl) if err != nil { diff --git a/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke_test.go b/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke_test.go index d6dd9748ba5..0273177d9df 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke_test.go +++ b/cli/azd/extensions/azure.ai.rle/internal/cmd/invoke_test.go @@ -20,12 +20,15 @@ import ( "slices" "strconv" "strings" + "sync/atomic" "testing" "time" + "azure.ai.rle/internal/project" "azure.ai.rle/internal/ui" "github.com/azure/azure-dev/cli/azd/pkg/azdext" + "github.com/gorilla/websocket" ) const testFoundryProjectPath = "/api/projects/project-1" @@ -51,6 +54,27 @@ func TestInvokeRemoteCreatesInstanceAndRunsShell(t *testing.T) { switch r.URL.Path { case "/health": _, _ = w.Write([]byte(`{"status":"healthy"}`)) + case "/ws": + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + var request map[string]any + if err := connection.ReadJSON(&request); err != nil { + t.Errorf("read WebSocket request: %v", err) + return + } + if request["type"] != "state" { + t.Errorf("expected state request, got %#v", request) + } + if err := connection.WriteJSON(map[string]any{ + "type": "state", + "data": map[string]any{"state": "ready"}, + }); err != nil { + t.Errorf("write WebSocket response: %v", err) + } default: http.NotFound(w, r) } @@ -91,7 +115,7 @@ func TestInvokeRemoteCreatesInstanceAndRunsShell(t *testing.T) { useTestProjectEndpoint(t, controlPlane.URL) command := newInvokeCommand() - command.SetIn(strings.NewReader("health\nexit\n")) + command.SetIn(strings.NewReader("state\nexit\n")) var output bytes.Buffer command.SetOut(&output) command.SetErr(&output) @@ -104,8 +128,8 @@ func TestInvokeRemoteCreatesInstanceAndRunsShell(t *testing.T) { if strings.Contains(output.String(), envServer.URL) { t.Fatalf("expected instance data-plane URL to remain hidden, got %s", output.String()) } - if !strings.Contains(output.String(), `"status": "healthy"`) { - t.Fatalf("expected remote shell health output, got %s", output.String()) + if !strings.Contains(output.String(), `"state": "ready"`) { + t.Fatalf("expected remote shell state output, got %s", output.String()) } if !instanceDeleted || !groupDeleted { t.Fatal("expected remote invoke to delete the instance and group") @@ -925,22 +949,40 @@ func TestRemotePlaygroundProxyForwardsToSandbox(t *testing.T) { requestCount := 0 envServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { requestCount++ - w.Header().Set("Content-Type", "application/json") - switch r.URL.Path { - case "/web": - http.NotFound(w, r) - case "/state": - _, _ = w.Write([]byte(`{"step_count":3}`)) - default: + if r.URL.Path != "/ws" { http.NotFound(w, r) + return + } + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + var request map[string]any + if err := connection.ReadJSON(&request); err != nil { + t.Errorf("read WebSocket request: %v", err) + return + } + if request["type"] != "state" { + t.Errorf("expected state request, got %#v", request) + } + if err := connection.WriteJSON(map[string]any{ + "type": "state", + "data": map[string]any{"step_count": 3}, + }); err != nil { + t.Errorf("write WebSocket response: %v", err) } })) defer envServer.Close() + runtimeSession := project.NewWebSocketRuntimeSession(envServer.URL, 30, nil) + defer runtimeSession.Close() playgroundUrl, stop, err := remotePlaygroundUrlWithAuthorizationProvider( t.Context(), envServer.URL, nil, + runtimeSession, ) if err != nil { t.Fatal(err) @@ -986,14 +1028,78 @@ func TestRemotePlaygroundProxyForwardsToSandbox(t *testing.T) { if err != nil { t.Fatal(err) } - if string(body) != `{"step_count":3}` { - t.Fatalf("expected proxied state body, got %s", body) + var state map[string]any + if err := json.Unmarshal(body, &state); err != nil || state["step_count"] != float64(3) { + t.Fatalf("expected proxied state body, got %s (err: %v)", body, err) } if requestCount != 1 { t.Fatalf("expected one authorized backend request, got %d", requestCount) } } +func TestRemotePlaygroundCancellationDoesNotFailSharedSession(t *testing.T) { + firstRequestReceived := make(chan struct{}) + releaseFirstResponse := make(chan struct{}) + var requestCount atomic.Int32 + envServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + for { + var request map[string]any + if err := connection.ReadJSON(&request); err != nil { + if websocket.IsCloseError(err, websocket.CloseNormalClosure) { + return + } + t.Errorf("read WebSocket request: %v", err) + return + } + currentRequest := requestCount.Add(1) + if currentRequest == 1 { + close(firstRequestReceived) + <-releaseFirstResponse + } + if err := connection.WriteJSON(map[string]any{ + "type": "state", + "data": map[string]any{"request": currentRequest}, + }); err != nil { + t.Errorf("write WebSocket response: %v", err) + return + } + } + })) + defer envServer.Close() + + runtimeSession := project.NewWebSocketRuntimeSession(envServer.URL, 30, nil) + defer runtimeSession.Close() + requestContext, cancelRequest := context.WithCancel(t.Context()) + request := httptest.NewRequest(http.MethodGet, "/state", nil).WithContext(requestContext) + recorder := httptest.NewRecorder() + proxyDone := make(chan struct{}) + go func() { + proxyStatefulOpenEnvOperation(recorder, request, "state", runtimeSession) + close(proxyDone) + }() + + <-firstRequestReceived + cancelRequest() + close(releaseFirstResponse) + <-proxyDone + + if recorder.Code != http.StatusOK { + t.Fatalf("expected canceled browser request to drain successfully, got %d", recorder.Code) + } + if _, err := runtimeSession.Call(t.Context(), "state", ""); err != nil { + t.Fatalf("expected shared session to remain usable: %v", err) + } + if requestCount.Load() != 2 { + t.Fatalf("expected two state requests on the shared session, got %d", requestCount.Load()) + } +} + func TestRemotePlaygroundProxyRefreshesAuthorizationForEachRequest(t *testing.T) { var authorizations []string envServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { diff --git a/cli/azd/extensions/azure.ai.rle/internal/cmd/root.go b/cli/azd/extensions/azure.ai.rle/internal/cmd/root.go index a650ca91f66..b17411e7bb7 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/cmd/root.go +++ b/cli/azd/extensions/azure.ai.rle/internal/cmd/root.go @@ -38,7 +38,7 @@ func NewRootCommand() *cobra.Command { userCommands := []*cobra.Command{ newListCommand(&extCtx.OutputFormat), newShowCommand(&extCtx.OutputFormat), - newInitCommand(), + newInitCommand(&extCtx.NoPrompt), newInvokeCommand(), newPublishCommand(), newRunCommand(), diff --git a/cli/azd/extensions/azure.ai.rle/internal/cmd/root_test.go b/cli/azd/extensions/azure.ai.rle/internal/cmd/root_test.go index d0797f8d801..6d2420a86dd 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/cmd/root_test.go +++ b/cli/azd/extensions/azure.ai.rle/internal/cmd/root_test.go @@ -5,11 +5,16 @@ package cmd import ( "bytes" + "context" "encoding/json" + "errors" "os" "path/filepath" + "slices" "strings" "testing" + + "github.com/azure/azure-dev/cli/azd/pkg/azdext" ) func TestNewRootCommandIncludesExpectedCommands(t *testing.T) { @@ -232,19 +237,24 @@ func TestLifecycleCommandsRejectPositionalArguments(t *testing.T) { t.Fatalf("expected init command to be registered: %v", err) } if err := initCommand.Args(initCommand, []string{"custom_env"}); err != nil { - t.Fatalf("expected init to accept one positional environment name: %v", err) + t.Fatalf("expected init to accept one positional folder name: %v", err) + } + if err := initCommand.Args(initCommand, nil); err != nil { + t.Fatalf("expected interactive init to allow an omitted folder name: %v", err) } if err := initCommand.Args(initCommand, []string{"one", "two"}); err == nil { t.Fatal("expected init to reject multiple positional arguments") } } -func TestInitCopiesOpenEnvEchoSampleByDefault(t *testing.T) { +func TestInitSelectsSampleAndCopiesItToNamedFolder(t *testing.T) { tempDir := t.TempDir() t.Chdir(tempDir) - stubOpenEnvEchoCheckout(t) + stubRleSampleCatalog(t, []string{"echo", "wordle"}, "wordle", "training_env") - command := newInitCommand() + noPrompt := false + command := newInitCommand(&noPrompt) + command.SetArgs([]string{"training_env"}) var output bytes.Buffer command.SetOut(&output) command.SetErr(&output) @@ -252,7 +262,7 @@ func TestInitCopiesOpenEnvEchoSampleByDefault(t *testing.T) { t.Fatal(err) } - sessionDir := filepath.Join(tempDir, "echo_env") + sessionDir := filepath.Join(tempDir, "training_env") // The test reads state from its own temporary session directory. stateBytes, err := os.ReadFile(filepath.Join(sessionDir, rleStateFile)) //nolint:gosec if err != nil { @@ -262,11 +272,11 @@ func TestInitCopiesOpenEnvEchoSampleByDefault(t *testing.T) { if err := json.Unmarshal(stateBytes, &state); err != nil { t.Fatal(err) } - if state.EnvironmentName != "echo_env" { - t.Fatalf("expected echo_env environment name, got %q", state.EnvironmentName) + if state.EnvironmentName != "training_env" { + t.Fatalf("expected training_env environment name, got %q", state.EnvironmentName) } if _, err := os.Stat(filepath.Join(sessionDir, "server", "Dockerfile")); err != nil { - t.Fatalf("expected copied OpenEnv server Dockerfile: %v", err) + t.Fatalf("expected copied RLE sample server Dockerfile: %v", err) } if _, err := os.Stat(filepath.Join(sessionDir, ".git")); !os.IsNotExist(err) { t.Fatalf("expected copied sample not to include .git metadata, got err=%v", err) @@ -274,38 +284,79 @@ func TestInitCopiesOpenEnvEchoSampleByDefault(t *testing.T) { if strings.Contains(output.String(), sessionDir) { t.Fatalf("expected init output not to use absolute cd path, got %s", output.String()) } - expectedCd := `cd "` + "." + string(os.PathSeparator) + "echo_env" + `"` + expectedCd := `cd "` + "." + string(os.PathSeparator) + "training_env" + `"` if !strings.Contains(output.String(), expectedCd) { t.Fatalf("expected init output to quote relative cd path, got %s", output.String()) } + if !strings.Contains(output.String(), `Copied RLE sample "wordle".`) { + t.Fatalf("expected selected sample in output, got %s", output.String()) + } } -func TestInitUsesPositionalNameForDefaultSample(t *testing.T) { +func TestInitWithoutFolderUsesSelectedSampleName(t *testing.T) { tempDir := t.TempDir() t.Chdir(tempDir) - stubOpenEnvEchoCheckout(t) + stubRleSampleCatalog(t, []string{"echo", "wordle"}, "wordle", "wordle") - command := newInitCommand() - command.SetArgs([]string{"code_rl"}) - var output bytes.Buffer - command.SetOut(&output) - command.SetErr(&output) + noPrompt := false + command := newInitCommand(&noPrompt) + command.SetArgs(nil) if err := command.Execute(); err != nil { t.Fatal(err) } + if _, err := os.Stat(filepath.Join(tempDir, "wordle", rleStateFile)); err != nil { + t.Fatalf("expected selected sample folder and state: %v", err) + } +} - sessionDir := filepath.Join(tempDir, "code_rl") - // The test reads state from its own temporary session directory. - stateBytes, err := os.ReadFile(filepath.Join(sessionDir, rleStateFile)) //nolint:gosec - if err != nil { - t.Fatal(err) +func TestInitNoPromptUsesPositionalNameAsSampleAndFolder(t *testing.T) { + tempDir := t.TempDir() + t.Chdir(tempDir) + t.Setenv(rleEnableEnvVar, "true") + + oldLoad := loadRleSampleCatalogFunc + oldSelect := selectRleSampleFunc + loadRleSampleCatalogFunc = func() (rleSampleCatalog, error) { + return &testRleSampleCatalog{ + t: t, + sampleNames: []string{"echo", "wordle"}, + expectedSampleName: "wordle", + expectedFolderName: "wordle", + }, nil + } + selectRleSampleFunc = func(context.Context, []string) (string, error) { + t.Fatal("expected no-prompt init to bypass the prompt") + return "", nil } - var state rleState - if err := json.Unmarshal(stateBytes, &state); err != nil { + t.Cleanup(func() { + loadRleSampleCatalogFunc = oldLoad + selectRleSampleFunc = oldSelect + }) + + command := NewRootCommand() + command.SetArgs([]string{"init", "wordle", "--no-prompt"}) + if err := command.Execute(); err != nil { t.Fatal(err) } - if state.EnvironmentName != "code_rl" { - t.Fatalf("expected code_rl environment name, got %q", state.EnvironmentName) +} + +func TestInitNoPromptRequiresPositionalSampleName(t *testing.T) { + t.Setenv(rleEnableEnvVar, "true") + command := NewRootCommand() + command.SetArgs([]string{"init", "--no-prompt"}) + err := command.Execute() + localError, ok := errors.AsType[*azdext.LocalError](err) + if !ok || localError.Code != "rle_sample_name_required" { + t.Fatalf("expected missing no-prompt sample error, got %v", err) + } +} + +func TestResolveRleSampleRejectsUnknownNonInteractiveSample(t *testing.T) { + _, err := resolveRleSample(t.Context(), "missing", []string{"echo", "wordle"}) + localError, ok := errors.AsType[*azdext.LocalError](err) + if !ok || localError.Code != "rle_sample_not_found" || + !strings.Contains(localError.Suggestion, "echo, wordle") { + t.Fatalf("expected available sample guidance, got %v", err) } } @@ -368,32 +419,74 @@ func TestInitNextStepsUseShellAppropriateSyntax(t *testing.T) { } } -func stubOpenEnvEchoCheckout(t *testing.T) { - t.Helper() - old := checkoutOpenEnvEchoSampleFunc - checkoutOpenEnvEchoSampleFunc = func(name string, dest string, force bool) (string, error) { - sessionDir := filepath.Join(dest, name) - if force { - if err := os.RemoveAll(sessionDir); err != nil { - return "", err - } - } - if err := os.MkdirAll(sessionDir, 0750); err != nil { - return "", err - } - serverDir := filepath.Join(sessionDir, "server") - if err := os.MkdirAll(serverDir, 0750); err != nil { - return "", err - } - if err := os.WriteFile(filepath.Join(serverDir, "Dockerfile"), []byte("FROM scratch\n"), 0600); err != nil { +type testRleSampleCatalog struct { + t *testing.T + sampleNames []string + expectedSampleName string + expectedFolderName string +} + +func (c *testRleSampleCatalog) SampleNames() []string { + return c.sampleNames +} + +func (c *testRleSampleCatalog) Copy( + sampleName string, + folderName string, + dest string, + force bool, +) (string, error) { + c.t.Helper() + if sampleName != c.expectedSampleName { + c.t.Fatalf("expected RLE sample %q, got %q", c.expectedSampleName, sampleName) + } + if folderName != c.expectedFolderName { + c.t.Fatalf("expected folder name %q, got %q", c.expectedFolderName, folderName) + } + sessionDir := filepath.Join(dest, folderName) + if force { + if err := os.RemoveAll(sessionDir); err != nil { return "", err } - if err := os.WriteFile(filepath.Join(sessionDir, "openenv.yaml"), []byte("name: echo_env\n"), 0600); err != nil { - return "", err + } + if err := os.MkdirAll(filepath.Join(sessionDir, "server"), 0750); err != nil { + return "", err + } + if err := os.WriteFile(filepath.Join(sessionDir, "server", "Dockerfile"), []byte("FROM scratch\n"), 0600); err != nil { + return "", err + } + return sessionDir, nil +} + +func (c *testRleSampleCatalog) Close() error { + return nil +} + +func stubRleSampleCatalog( + t *testing.T, + sampleNames []string, + selectedSampleName string, + expectedFolderName string, +) { + t.Helper() + oldLoad := loadRleSampleCatalogFunc + oldSelect := selectRleSampleFunc + loadRleSampleCatalogFunc = func() (rleSampleCatalog, error) { + return &testRleSampleCatalog{ + t: t, + sampleNames: sampleNames, + expectedSampleName: selectedSampleName, + expectedFolderName: expectedFolderName, + }, nil + } + selectRleSampleFunc = func(_ context.Context, actualSampleNames []string) (string, error) { + if !slices.Equal(actualSampleNames, sampleNames) { + t.Fatalf("expected sample names %v, got %v", sampleNames, actualSampleNames) } - return sessionDir, nil + return selectedSampleName, nil } t.Cleanup(func() { - checkoutOpenEnvEchoSampleFunc = old + loadRleSampleCatalogFunc = oldLoad + selectRleSampleFunc = oldSelect }) } diff --git a/cli/azd/extensions/azure.ai.rle/internal/project/runtime.go b/cli/azd/extensions/azure.ai.rle/internal/project/runtime.go index d3017495189..75bf1cc72c3 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/project/runtime.go +++ b/cli/azd/extensions/azure.ai.rle/internal/project/runtime.go @@ -26,6 +26,8 @@ type callOptions struct { type AuthorizationProvider func(context.Context) (string, error) +type runtimeCaller func(context.Context, string, string, *callOptions, AuthorizationProvider) (string, error) + func RunShellWithContext( ctx context.Context, input io.Reader, @@ -80,6 +82,18 @@ func runShell( baseUrl string, timeout int, authorizationProvider AuthorizationProvider, +) error { + return runShellWithCaller(ctx, input, output, baseUrl, timeout, authorizationProvider, call) +} + +func runShellWithCaller( + ctx context.Context, + input io.Reader, + output io.Writer, + baseUrl string, + timeout int, + authorizationProvider AuthorizationProvider, + caller runtimeCaller, ) error { fmt.Fprintln(output, "Environment runtime shell. Type help for commands, exit to quit.") scanner := bufio.NewScanner(input) @@ -125,7 +139,7 @@ func runShell( } } - response, err := call(ctx, baseUrl, operation, flags, authorizationProvider) + response, err := caller(ctx, baseUrl, operation, flags, authorizationProvider) if err != nil { fmt.Fprintf(output, "error: %v\n", err) continue diff --git a/cli/azd/extensions/azure.ai.rle/internal/project/scaffold.go b/cli/azd/extensions/azure.ai.rle/internal/project/scaffold.go index c0a1a3b3a67..6cd5a6ee91a 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/project/scaffold.go +++ b/cli/azd/extensions/azure.ai.rle/internal/project/scaffold.go @@ -7,17 +7,22 @@ import ( "os" "os/exec" "path/filepath" + "slices" "strings" "github.com/azure/azure-dev/cli/azd/pkg/azdext" ) const ( - openEnvRepoUrl = "https://github.com/huggingface/OpenEnv.git" - openEnvRepoRef = "main" - openEnvEchoSamplePath = "envs/echo_env" + rleSamplesRepoURL = "https://github.com/sujit-kamireddy/rle-samples.git" + rleSamplesRepoRef = "main" ) +type RleSampleCatalog struct { + repoDir string + sampleNames []string +} + func createRleSessionDir(name string, dest string, force bool) (string, error) { sessionDir := filepath.Join(dest, name) if entries, err := os.ReadDir(sessionDir); err == nil && len(entries) > 0 && !force { @@ -41,66 +46,138 @@ func createRleSessionDir(name string, dest string, force bool) (string, error) { return sessionDir, nil } -func CheckoutOpenEnvEchoSample(name string, dest string, force bool) (string, error) { - name, err := ValidateEnvironmentName(name) - if err != nil { - return "", err - } - sessionDir, err := createRleSessionDir(name, dest, force) - if err != nil { - return "", err - } - tempDir, err := os.MkdirTemp("", "azd-rle-open-env-*") +func LoadRleSampleCatalog() (*RleSampleCatalog, error) { + return loadRleSampleCatalog(rleSamplesRepoURL, rleSamplesRepoRef) +} + +func loadRleSampleCatalog(repoURL string, repoRef string) (*RleSampleCatalog, error) { + tempDir, err := os.MkdirTemp("", "azd-rle-samples-*") if err != nil { - return "", err + return nil, err } - defer func() { - _ = os.RemoveAll(tempDir) - }() - - if err := runGitCheckout( + if _, err := runGitCommand( "clone", "--depth", "1", "--filter=blob:none", "--sparse", - "--branch", openEnvRepoRef, - openEnvRepoUrl, + "--branch", repoRef, + "--single-branch", + repoURL, tempDir, ); err != nil { + _ = os.RemoveAll(tempDir) + return nil, err + } + + sampleNames, err := listRleSamples(tempDir, repoRef) + if err != nil { + _ = os.RemoveAll(tempDir) + return nil, err + } + if len(sampleNames) == 0 { + _ = os.RemoveAll(tempDir) + return nil, &azdext.LocalError{ + Message: "The RLE samples repository does not contain any sample environments.", + Code: "rle_samples_empty", + Category: azdext.LocalErrorCategoryUser, + Suggestion: fmt.Sprintf( + "Add sample directories to %s, then retry.", + strings.TrimSuffix(rleSamplesRepoURL, ".git"), + ), + } + } + return &RleSampleCatalog{ + repoDir: tempDir, + sampleNames: sampleNames, + }, nil +} + +func (c *RleSampleCatalog) SampleNames() []string { + return slices.Clone(c.sampleNames) +} + +func (c *RleSampleCatalog) Copy(sampleName string, folderName string, dest string, force bool) (string, error) { + folderName, err := ValidateEnvironmentName(folderName) + if err != nil { return "", err } - if err := runGitCheckout("-C", tempDir, "sparse-checkout", "set", openEnvEchoSamplePath); err != nil { + sourcePath := filepath.ToSlash(filepath.Join("envs", sampleName)) + if _, err := runGitCommand("-C", c.repoDir, "sparse-checkout", "set", sourcePath); err != nil { return "", err } + sourceDir := filepath.Join(c.repoDir, filepath.FromSlash(sourcePath)) + return copyRleSample(sourceDir, folderName, dest, force) +} + +func (c *RleSampleCatalog) Close() error { + return os.RemoveAll(c.repoDir) +} + +func listRleSamples(repoDir string, repoRef string) ([]string, error) { + output, err := runGitCommand( + "-C", + repoDir, + "ls-tree", + "-d", + "--name-only", + repoRef+":envs", + ) + if err != nil { + return nil, err + } + sampleNames := strings.Fields(string(output)) + sampleNames = slices.DeleteFunc(sampleNames, func(name string) bool { + return strings.HasPrefix(name, ".") + }) + slices.Sort(sampleNames) + return sampleNames, nil +} - sourceDir := filepath.Join(tempDir, filepath.FromSlash(openEnvEchoSamplePath)) +func copyRleSample(sourceDir string, folderName string, dest string, force bool) (string, error) { + sourceInfo, err := os.Stat(sourceDir) + if os.IsNotExist(err) { + return "", &azdext.LocalError{ + Message: fmt.Sprintf("RLE sample source %q was not found.", sourceDir), + Code: "rle_sample_source_not_found", + Category: azdext.LocalErrorCategoryInternal, + Suggestion: "Run azd ai rle init again to refresh the sample list.", + } + } else if err != nil { + return "", err + } else if !sourceInfo.IsDir() { + return "", fmt.Errorf("RLE sample source %q is not a directory", sourceDir) + } + sessionDir, err := createRleSessionDir(folderName, dest, force) + if err != nil { + return "", err + } if err := copyDirectory(sourceDir, sessionDir); err != nil { return "", err } return sessionDir, nil } -func runGitCheckout(args ...string) error { +func runGitCommand(args ...string) ([]byte, error) { if _, err := exec.LookPath("git"); err != nil { - return &azdext.LocalError{ + return nil, &azdext.LocalError{ Message: "Could not find \"git\" on PATH.", Code: "rle_git_not_found", Category: azdext.LocalErrorCategoryUser, Suggestion: "Install Git, then retry azd ai rle init.", } } - process := exec.Command("git", args...) //nolint:gosec // args are fixed by init's OpenEnv sample checkout flow. + process := exec.Command("git", args...) //nolint:gosec process.Env = os.Environ() output, err := process.CombinedOutput() if err != nil { - return &azdext.LocalError{ - Message: fmt.Sprintf("Failed to checkout OpenEnv echo sample: %v", err), - Code: "rle_open_env_checkout_failed", + return nil, &azdext.LocalError{ + Message: fmt.Sprintf("Failed to download RLE samples: %v", err), + Code: "rle_samples_download_failed", Category: azdext.LocalErrorCategoryUser, Suggestion: strings.TrimSpace(string(output)), } } - return nil + return output, nil } func copyDirectory(sourceDir string, destDir string) error { diff --git a/cli/azd/extensions/azure.ai.rle/internal/project/scaffold_test.go b/cli/azd/extensions/azure.ai.rle/internal/project/scaffold_test.go index 23d701ee1ca..da63f32549c 100644 --- a/cli/azd/extensions/azure.ai.rle/internal/project/scaffold_test.go +++ b/cli/azd/extensions/azure.ai.rle/internal/project/scaffold_test.go @@ -5,8 +5,10 @@ package project import ( "os" + "os/exec" "path/filepath" "runtime" + "slices" "testing" ) @@ -71,17 +73,105 @@ func TestCopyDirectoryRejectsFileSource(t *testing.T) { } } -func TestCheckoutOpenEnvEchoSampleRejectsInvalidNameBeforeChangingDestination(t *testing.T) { +func TestRleSampleCatalogUsesSparseCheckout(t *testing.T) { + if _, err := exec.LookPath("git"); err != nil { + t.Skip("git is not available") + } + + sourceRepo := t.TempDir() + runTestGit(t, sourceRepo, "init", "--initial-branch=main") + for _, sampleName := range []string{"code_rl", "math_rl"} { + sampleDir := filepath.Join(sourceRepo, "envs", sampleName) + if err := os.MkdirAll(sampleDir, 0750); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(sampleDir, "sample.txt"), []byte(sampleName), 0600); err != nil { + t.Fatal(err) + } + } + if err := os.WriteFile(filepath.Join(sourceRepo, "README.md"), []byte("samples"), 0600); err != nil { + t.Fatal(err) + } + runTestGit(t, sourceRepo, "add", ".") + runTestGit( + t, + sourceRepo, + "-c", "user.name=RLE Tests", + "-c", "user.email=rle-tests@example.com", + "commit", "-m", "Add samples", + ) + + catalog, err := loadRleSampleCatalog(sourceRepo, "main") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := catalog.Close(); err != nil { + t.Errorf("close sample catalog: %v", err) + } + }) + if !slices.Equal(catalog.SampleNames(), []string{"code_rl", "math_rl"}) { + t.Fatalf("expected sorted sample names, got %v", catalog.SampleNames()) + } + if _, err := os.Stat(filepath.Join(catalog.repoDir, "envs")); !os.IsNotExist(err) { + t.Fatalf("expected sample contents not to be checked out before selection, got err=%v", err) + } + + sessionDir, err := catalog.Copy("math_rl", "training_env", t.TempDir(), false) + if err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(sessionDir, "sample.txt")); err != nil { + t.Fatalf("expected selected sample to be copied: %v", err) + } + if _, err := os.Stat(filepath.Join(catalog.repoDir, "envs", "code_rl")); !os.IsNotExist(err) { + t.Fatalf("expected unselected sample not to be checked out, got err=%v", err) + } +} + +func TestCopyRleSampleRenamesDestination(t *testing.T) { + sourceDir := t.TempDir() + if err := os.WriteFile(filepath.Join(sourceDir, "sample.txt"), []byte("content"), 0600); err != nil { + t.Fatal(err) + } destDir := t.TempDir() - sentinel := filepath.Join(destDir, "sentinel.txt") + sessionDir, err := copyRleSample(sourceDir, "my_environment", destDir, false) + if err != nil { + t.Fatal(err) + } + if sessionDir != filepath.Join(destDir, "my_environment") { + t.Fatalf("expected renamed destination, got %q", sessionDir) + } + if _, err := os.Stat(filepath.Join(sessionDir, "sample.txt")); err != nil { + t.Fatalf("expected sample file in renamed destination: %v", err) + } +} + +func runTestGit(t *testing.T, dir string, args ...string) { + t.Helper() + command := exec.Command("git", args...) //nolint:gosec + command.Dir = dir + if output, err := command.CombinedOutput(); err != nil { + t.Fatalf("git %v failed: %v\n%s", args, err, output) + } +} + +func TestCopyRleSampleValidatesSourceBeforeReplacingDestination(t *testing.T) { + destDir := t.TempDir() + sessionDir := filepath.Join(destDir, "my_environment") + if err := os.MkdirAll(sessionDir, 0750); err != nil { + t.Fatal(err) + } + sentinel := filepath.Join(sessionDir, "keep.txt") if err := os.WriteFile(sentinel, []byte("keep"), 0600); err != nil { t.Fatal(err) } - if _, err := CheckoutOpenEnvEchoSample("../bad", destDir, true); err == nil { - t.Fatal("expected invalid environment name to be rejected") + _, err := copyRleSample(filepath.Join(t.TempDir(), "missing"), "my_environment", destDir, true) + if err == nil { + t.Fatal("expected missing RLE sample to fail") } - if _, err := os.Stat(sentinel); err != nil { - t.Fatalf("expected destination to be unchanged: %v", err) + if _, statErr := os.Stat(sentinel); statErr != nil { + t.Fatalf("expected destination to remain unchanged after sample lookup failure: %v", statErr) } } diff --git a/cli/azd/extensions/azure.ai.rle/internal/project/websocket_runtime.go b/cli/azd/extensions/azure.ai.rle/internal/project/websocket_runtime.go new file mode 100644 index 00000000000..f69972bd384 --- /dev/null +++ b/cli/azd/extensions/azure.ai.rle/internal/project/websocket_runtime.go @@ -0,0 +1,566 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package project + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "strings" + "sync" + "time" + + "github.com/azure/azure-dev/cli/azd/pkg/azdext" + "github.com/gorilla/websocket" +) + +const ( + maxWebSocketMessageBytes = 8 * 1024 * 1024 + webSocketHandshakeTimeout = 30 * time.Second + webSocketPingInterval = 20 * time.Second + webSocketPingTimeout = 20 * time.Second + webSocketDrainTimeout = 60 * time.Second +) + +var defaultWebSocketHandshakeRetryDelays = []time.Duration{time.Second, 2 * time.Second} + +type WebSocketRuntimeSession struct { + baseURL string + timeout int + authorizationProvider AuthorizationProvider + mu sync.Mutex + exchangeMu sync.Mutex + connection *websocket.Conn + connectionDone chan struct{} + keepAliveInterval time.Duration + handshakeRetryDelays []time.Duration + drainTimeout time.Duration + terminalError error + closed bool +} + +func NewWebSocketRuntimeSession( + baseURL string, + timeout int, + authorizationProvider AuthorizationProvider, +) *WebSocketRuntimeSession { + return &WebSocketRuntimeSession{ + baseURL: baseURL, + timeout: timeout, + authorizationProvider: authorizationProvider, + keepAliveInterval: webSocketPingInterval, + handshakeRetryDelays: defaultWebSocketHandshakeRetryDelays, + drainTimeout: webSocketDrainTimeout, + } +} + +// RunWebSocketShellWithContextAndAuthorizationProvider uses one persistent WebSocket +// for stateful OpenEnv operations while retaining the safe HTTP endpoints. +func RunWebSocketShellWithContextAndAuthorizationProvider( + ctx context.Context, + input io.Reader, + output io.Writer, + baseURL string, + timeout int, + authorizationProvider AuthorizationProvider, +) error { + session := NewWebSocketRuntimeSession(baseURL, timeout, authorizationProvider) + defer session.Close() + return RunWebSocketShellWithSession(ctx, input, output, session) +} + +func RunWebSocketShellWithSession( + ctx context.Context, + input io.Reader, + output io.Writer, + session *WebSocketRuntimeSession, +) error { + done := make(chan error, 1) + go func() { + done <- runShellWithCaller( + ctx, + input, + output, + session.baseURL, + session.timeout, + session.authorizationProvider, + session.call, + ) + }() + select { + case err := <-done: + return err + case <-ctx.Done(): + session.Close() + fmt.Fprintln(output) + return nil + } +} + +func (c *WebSocketRuntimeSession) call( + ctx context.Context, + baseURL string, + operation string, + flags *callOptions, + authorizationProvider AuthorizationProvider, +) (string, error) { + switch operation { + case "reset", "step", "state": + return c.exchange(ctx, operation, flags, true) + default: + return call(ctx, baseURL, operation, flags, authorizationProvider) + } +} + +func (c *WebSocketRuntimeSession) Call( + ctx context.Context, + operation string, + payload string, +) (string, error) { + flags := &callOptions{timeout: c.timeout} + switch operation { + case "reset": + flags.body = payload + case "step": + flags.action = payload + } + return c.call(ctx, c.baseURL, operation, flags, c.authorizationProvider) +} + +// CallAndDrain preserves cancellation until a request is sent, then drains its response +// so a disconnected HTTP client cannot disrupt the shared WebSocket protocol sequence. +func (c *WebSocketRuntimeSession) CallAndDrain( + ctx context.Context, + operation string, + payload string, +) (string, error) { + flags := &callOptions{timeout: c.timeout} + switch operation { + case "reset": + flags.body = payload + case "step": + flags.action = payload + } + return c.exchange(ctx, operation, flags, false) +} + +func (c *WebSocketRuntimeSession) exchange( + ctx context.Context, + operation string, + flags *callOptions, + cancelAfterSend bool, +) (string, error) { + c.exchangeMu.Lock() + defer c.exchangeMu.Unlock() + deadline, hasDeadline := operationDeadline(ctx, c.timeout) + operationCtx := ctx + cancel := func() {} + if hasDeadline { + operationCtx, cancel = context.WithDeadline(ctx, deadline) + } + defer cancel() + if err := operationCtx.Err(); err != nil { + return "", err + } + if err := c.connect(operationCtx); err != nil { + return "", err + } + c.mu.Lock() + connection := c.connection + terminalError := c.terminalError + c.mu.Unlock() + if connection == nil { + if terminalError != nil { + return "", terminalError + } + return "", fmt.Errorf("OpenEnv WebSocket session closed") + } + + request, err := webSocketRequest(operation, flags) + if err != nil { + return "", err + } + if len(request) > maxWebSocketMessageBytes { + return "", &azdext.LocalError{ + Message: fmt.Sprintf( + "OpenEnv WebSocket %s request is %d bytes; the RLE service limit is %d bytes.", + operation, + len(request), + maxWebSocketMessageBytes, + ), + Code: "rle_open_env_websocket_request_too_large", + Category: azdext.LocalErrorCategoryUser, + Suggestion: "Reduce the request payload and retry.", + } + } + + if hasDeadline { + if err := connection.SetWriteDeadline(deadline); err != nil { + return "", c.failConnection(connection, fmt.Errorf("set OpenEnv WebSocket write deadline: %w", err)) + } + if err := connection.SetReadDeadline(deadline); err != nil { + return "", c.failConnection(connection, fmt.Errorf("set OpenEnv WebSocket read deadline: %w", err)) + } + } else { + _ = connection.SetWriteDeadline(time.Time{}) + _ = connection.SetReadDeadline(time.Time{}) + } + + if err := connection.WriteMessage(websocket.TextMessage, request); err != nil { + return "", c.failConnection( + connection, + fmt.Errorf("send OpenEnv WebSocket %s request: %w", operation, err), + ) + } + exchangeDone := make(chan struct{}) + cancellationHandled := make(chan struct{}) + go func() { + select { + case <-exchangeDone: + case <-ctx.Done(): + if cancelAfterSend { + _ = c.failConnection( + connection, + fmt.Errorf("OpenEnv WebSocket %s request canceled: %w", operation, ctx.Err()), + ) + } else { + drainDeadline := time.Now().Add(c.drainTimeout) + if hasDeadline && deadline.Before(drainDeadline) { + drainDeadline = deadline + } + _ = connection.SetReadDeadline(drainDeadline) + } + } + close(cancellationHandled) + }() + messageType, response, err := connection.ReadMessage() + close(exchangeDone) + <-cancellationHandled + if err != nil { + return "", c.failConnection( + connection, + fmt.Errorf("receive OpenEnv WebSocket %s response: %w", operation, err), + ) + } + if messageType != websocket.TextMessage { + err := &azdext.LocalError{ + Message: "Environment runtime returned a non-text WebSocket response.", + Code: "rle_open_env_websocket_protocol_error", + Category: azdext.LocalErrorCategoryInternal, + } + return "", c.failConnection(connection, err) + } + result, terminal, err := parseWebSocketResponse(operation, response) + if terminal { + return "", c.failConnection(connection, err) + } + return result, err +} + +func (c *WebSocketRuntimeSession) connect(ctx context.Context) error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return c.terminalError + } + if c.terminalError != nil { + return c.terminalError + } + if c.connection != nil { + return nil + } + + connectCtx := ctx + cancel := func() {} + if deadline, hasDeadline := operationDeadline(ctx, c.timeout); hasDeadline { + connectCtx, cancel = context.WithDeadline(ctx, deadline) + } + defer cancel() + + endpoint, err := RuntimeWebSocketURL(c.baseURL) + if err != nil { + return err + } + headers := http.Header{} + if c.authorizationProvider != nil { + authorization, err := c.authorizationProvider(ctx) + if err != nil { + return fmt.Errorf("authenticate to environment runtime: %w", err) + } + if authorization != "" { + headers.Set("Authorization", authorization) + } + } + + dialer := *websocket.DefaultDialer + dialer.HandshakeTimeout = webSocketHandshakeTimeout + if c.timeout > 0 && time.Duration(c.timeout)*time.Second < dialer.HandshakeTimeout { + dialer.HandshakeTimeout = time.Duration(c.timeout) * time.Second + } + + for attempt := 0; ; attempt++ { + connection, response, err := dialer.DialContext(connectCtx, endpoint, headers) + if err == nil { + connection.SetReadLimit(maxWebSocketMessageBytes) + c.connection = connection + c.connectionDone = make(chan struct{}) + go c.keepAlive(connection, c.connectionDone) + return nil + } + detail := "" + retryable := isRetryableWebSocketHandshakeError(err) + if response != nil { + retryable = isRetryableWebSocketHandshakeStatus(response.StatusCode) + detail = readHealthErrorDetail(response.Body) + _ = response.Body.Close() + } + connectionError := &azdext.LocalError{ + Message: fmt.Sprintf( + "Environment runtime WebSocket connection failed%s: %v", + detail, + err, + ), + Code: "rle_open_env_websocket_connection_failed", + Category: azdext.LocalErrorCategoryUser, + Suggestion: "Check the remote RLE instance status and retry invoke.", + } + if !retryable || attempt >= len(c.handshakeRetryDelays) { + return connectionError + } + timer := time.NewTimer(c.handshakeRetryDelays[attempt]) + select { + case <-timer.C: + case <-connectCtx.Done(): + if !timer.Stop() { + <-timer.C + } + return connectCtx.Err() + } + } +} + +func isRetryableWebSocketHandshakeError(err error) bool { + _, ok := errors.AsType[net.Error](err) + return ok +} + +func isRetryableWebSocketHandshakeStatus(statusCode int) bool { + switch statusCode { + case http.StatusRequestTimeout, + http.StatusTooManyRequests, + http.StatusInternalServerError, + http.StatusBadGateway, + http.StatusServiceUnavailable, + http.StatusGatewayTimeout: + return true + default: + return false + } +} + +func (c *WebSocketRuntimeSession) keepAlive(connection *websocket.Conn, done <-chan struct{}) { + ticker := time.NewTicker(c.keepAliveInterval) + defer ticker.Stop() + for { + select { + case <-done: + return + case <-ticker.C: + deadline := time.Now().Add(webSocketPingTimeout) + if err := connection.WriteControl(websocket.PingMessage, nil, deadline); err != nil { + _ = c.failConnection(connection, fmt.Errorf("send OpenEnv WebSocket keepalive: %w", err)) + return + } + } + } +} + +func (c *WebSocketRuntimeSession) failConnection(connection *websocket.Conn, err error) error { + c.mu.Lock() + if c.connection == connection { + c.connection = nil + c.stopKeepAliveLocked() + } + if c.terminalError == nil { + c.terminalError = &azdext.LocalError{ + Message: fmt.Sprintf("The OpenEnv WebSocket session is no longer usable: %v", err), + Code: "rle_open_env_websocket_session_failed", + Category: azdext.LocalErrorCategoryUser, + Suggestion: "Exit and run invoke again to start a new environment session.", + } + } + terminalError := c.terminalError + c.mu.Unlock() + _ = connection.Close() + return terminalError +} + +func (c *WebSocketRuntimeSession) Close() { + c.mu.Lock() + connection := c.connection + c.connection = nil + c.stopKeepAliveLocked() + c.closed = true + if c.terminalError == nil { + c.terminalError = &azdext.LocalError{ + Message: "The OpenEnv WebSocket session is closed.", + Code: "rle_open_env_websocket_session_closed", + Category: azdext.LocalErrorCategoryUser, + Suggestion: "Run invoke again to start a new environment session.", + } + } + c.mu.Unlock() + if connection == nil { + return + } + deadline := time.Now().Add(2 * time.Second) + _ = connection.WriteControl( + websocket.CloseMessage, + websocket.FormatCloseMessage(websocket.CloseNormalClosure, "RLE invoke complete."), + deadline, + ) + _ = connection.Close() +} + +func (c *WebSocketRuntimeSession) stopKeepAliveLocked() { + if c.connectionDone != nil { + close(c.connectionDone) + c.connectionDone = nil + } +} + +func parseWebSocketResponse(operation string, response []byte) (string, bool, error) { + var envelope struct { + Type string `json:"type"` + Data json.RawMessage `json:"data"` + } + if err := json.Unmarshal(response, &envelope); err != nil { + return "", true, &azdext.LocalError{ + Message: fmt.Sprintf("Environment runtime returned invalid WebSocket JSON: %v", err), + Code: "rle_open_env_websocket_protocol_error", + Category: azdext.LocalErrorCategoryInternal, + } + } + if envelope.Type == "error" { + var errorData struct { + Code string `json:"code"` + } + _ = json.Unmarshal(envelope.Data, &errorData) + terminal := isTerminalOpenEnvError(errorData.Code) + suggestion := "Check the request payload and retry." + if terminal { + suggestion = "Exit and run invoke again to start a new environment session." + } + return "", terminal, &azdext.LocalError{ + Message: fmt.Sprintf( + "Environment runtime rejected the %s request: %s", + operation, + prettyJson(envelope.Data), + ), + Code: "rle_open_env_websocket_request_failed", + Category: azdext.LocalErrorCategoryUser, + Suggestion: suggestion, + } + } + expectedType := "observation" + if operation == "state" { + expectedType = "state" + } + if envelope.Type != expectedType || len(envelope.Data) == 0 { + return "", true, &azdext.LocalError{ + Message: fmt.Sprintf( + "Environment runtime returned WebSocket response type %q for %s; expected %q.", + envelope.Type, + operation, + expectedType, + ), + Code: "rle_open_env_websocket_protocol_error", + Category: azdext.LocalErrorCategoryInternal, + } + } + return prettyJson(envelope.Data), false, nil +} + +func isTerminalOpenEnvError(code string) bool { + switch code { + case "CAPACITY_REACHED", "FACTORY_ERROR", "SESSION_ERROR": + return true + default: + return false + } +} + +func RuntimeWebSocketURL(baseURL string) (string, error) { + endpoint, err := url.Parse(baseURL) + if err != nil { + return "", fmt.Errorf("parse environment runtime URL: %w", err) + } + switch strings.ToLower(endpoint.Scheme) { + case "http": + endpoint.Scheme = "ws" + case "https": + endpoint.Scheme = "wss" + default: + return "", fmt.Errorf("environment runtime URL must use HTTP or HTTPS") + } + endpoint.Path = strings.TrimRight(endpoint.Path, "/") + "/ws" + endpoint.RawPath = "" + return endpoint.String(), nil +} + +func webSocketRequest(operation string, flags *callOptions) ([]byte, error) { + request := map[string]any{"type": operation} + switch operation { + case "reset": + data, err := webSocketData(flags.body, "body", true) + if err != nil { + return nil, err + } + request["data"] = data + case "step": + data, err := webSocketData(flags.action, "action", false) + if err != nil { + return nil, err + } + request["data"] = data + case "state": + default: + return nil, fmt.Errorf("operation %q is not supported over OpenEnv WebSocket", operation) + } + return json.Marshal(request) +} + +func webSocketData(value string, flagName string, allowEmpty bool) (map[string]any, error) { + if strings.TrimSpace(value) == "" && allowEmpty { + return map[string]any{}, nil + } + data, err := validateJsonObject(value, flagName) + if err != nil { + return nil, err + } + var decoded map[string]any + if err := json.Unmarshal(data, &decoded); err != nil { + return nil, err + } + return decoded, nil +} + +func operationDeadline(ctx context.Context, timeoutSeconds int) (time.Time, bool) { + var deadline time.Time + hasDeadline := false + if timeoutSeconds > 0 { + deadline = time.Now().Add(time.Duration(timeoutSeconds) * time.Second) + hasDeadline = true + } + if contextDeadline, ok := ctx.Deadline(); ok && (!hasDeadline || contextDeadline.Before(deadline)) { + deadline = contextDeadline + hasDeadline = true + } + return deadline, hasDeadline +} diff --git a/cli/azd/extensions/azure.ai.rle/internal/project/websocket_runtime_test.go b/cli/azd/extensions/azure.ai.rle/internal/project/websocket_runtime_test.go new file mode 100644 index 00000000000..a850dae4341 --- /dev/null +++ b/cli/azd/extensions/azure.ai.rle/internal/project/websocket_runtime_test.go @@ -0,0 +1,683 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package project + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "net" + "net/http" + "net/http/httptest" + "slices" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/azure/azure-dev/cli/azd/pkg/azdext" + "github.com/gorilla/websocket" +) + +func TestRunWebSocketShellUsesPersistentSocketForStatefulOperations(t *testing.T) { + var requests []map[string]any + var safePaths []string + upgrades := 0 + var captureMu sync.Mutex + upgrader := websocket.Upgrader{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Query().Get("api-version") != "test-version" { + t.Errorf("expected preserved API version, got %q", r.URL.Query().Get("api-version")) + } + if r.Header.Get("Authorization") != "******" { + t.Errorf("unexpected authorization header %q", r.Header.Get("Authorization")) + } + if r.URL.Path != "/ws" { + captureMu.Lock() + safePaths = append(safePaths, r.URL.Path) + captureMu.Unlock() + _, _ = fmt.Fprintf(w, `{"path":%q}`, r.URL.Path) //nolint:gosec // Test response uses JSON encoding. + return + } + + captureMu.Lock() + upgrades++ + captureMu.Unlock() + connection, err := upgrader.Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + for { + var request map[string]any + if err := connection.ReadJSON(&request); err != nil { + return + } + captureMu.Lock() + requests = append(requests, request) + captureMu.Unlock() + responseType := "observation" + if request["type"] == "state" { + responseType = "state" + } + if err := connection.WriteJSON(map[string]any{ + "type": responseType, + "data": map[string]any{"requestType": request["type"]}, + }); err != nil { + t.Errorf("write WebSocket response: %v", err) + return + } + } + })) + defer server.Close() + + authorizationCalls := 0 + authorizationProvider := func(context.Context) (string, error) { + authorizationCalls++ + return "******", nil + } + input := strings.NewReader( + "reset {\"seed\":42}\nstep {\"message\":\"hello\"}\nstate\nhealth\nmetadata\nschema\nexit\n", + ) + var output bytes.Buffer + err := RunWebSocketShellWithContextAndAuthorizationProvider( + t.Context(), + input, + &output, + server.URL+"?api-version=test-version", + 30, + authorizationProvider, + ) + if err != nil { + t.Fatal(err) + } + + captureMu.Lock() + capturedUpgrades := upgrades + capturedRequests := slices.Clone(requests) + capturedSafePaths := slices.Clone(safePaths) + captureMu.Unlock() + if capturedUpgrades != 1 { + t.Fatalf("expected one persistent WebSocket, got %d", capturedUpgrades) + } + if len(capturedRequests) != 3 { + t.Fatalf("expected three WebSocket requests, got %#v", capturedRequests) + } + assertWebSocketRequest(t, capturedRequests[0], "reset", map[string]any{"seed": float64(42)}) + assertWebSocketRequest(t, capturedRequests[1], "step", map[string]any{"message": "hello"}) + assertWebSocketRequest(t, capturedRequests[2], "state", nil) + if !slices.Equal(capturedSafePaths, []string{"/health", "/metadata", "/schema"}) { + t.Fatalf("unexpected safe HTTP operations: %v", capturedSafePaths) + } + if authorizationCalls != 4 { + t.Fatalf("expected one WebSocket and three HTTP authorization calls, got %d", authorizationCalls) + } + if !strings.Contains(output.String(), `"requestType": "step"`) { + t.Fatalf("expected formatted WebSocket response, got %s", output.String()) + } +} + +func TestWebSocketHandshakeRetriesTransientFailures(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if attempts.Add(1) < 3 { + http.Error(w, "temporarily unavailable", http.StatusServiceUnavailable) + return + } + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + if _, _, err := connection.ReadMessage(); err != nil { + return + } + if err := connection.WriteJSON(map[string]any{"type": "state", "data": map[string]any{}}); err != nil { + t.Errorf("write WebSocket response: %v", err) + } + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + session.handshakeRetryDelays = []time.Duration{time.Millisecond, time.Millisecond} + defer session.Close() + if _, err := session.Call(t.Context(), "state", ""); err != nil { + t.Fatal(err) + } + if attempts.Load() != 3 { + t.Fatalf("expected three handshake attempts, got %d", attempts.Load()) + } +} + +func TestWebSocketHandshakeDoesNotRetryPermanentFailure(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + http.Error(w, "unauthorized", http.StatusUnauthorized) + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + session.handshakeRetryDelays = []time.Duration{time.Millisecond, time.Millisecond} + defer session.Close() + if _, err := session.Call(t.Context(), "state", ""); err == nil { + t.Fatal("expected WebSocket handshake to fail") + } + if attempts.Load() != 1 { + t.Fatalf("expected one handshake attempt, got %d", attempts.Load()) + } +} + +func TestWebSocketHandshakeStopsAfterRetryLimit(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts.Add(1) + http.Error(w, "temporarily unavailable", http.StatusServiceUnavailable) + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + session.handshakeRetryDelays = []time.Duration{time.Millisecond, time.Millisecond} + defer session.Close() + if _, err := session.Call(t.Context(), "state", ""); err == nil { + t.Fatal("expected WebSocket handshake to fail") + } + if attempts.Load() != 3 { + t.Fatalf("expected three handshake attempts, got %d", attempts.Load()) + } +} + +func TestWebSocketHandshakeCancellationStopsBackoff(t *testing.T) { + attemptReceived := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case <-attemptReceived: + default: + close(attemptReceived) + } + http.Error(w, "temporarily unavailable", http.StatusServiceUnavailable) + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + session.handshakeRetryDelays = []time.Duration{time.Hour, time.Hour} + defer session.Close() + ctx, cancel := context.WithCancel(t.Context()) + callDone := make(chan error, 1) + go func() { + _, err := session.Call(ctx, "state", "") + callDone <- err + }() + <-attemptReceived + cancel() + select { + case err := <-callDone: + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context cancellation, got %v", err) + } + case <-time.After(time.Second): + t.Fatal("expected cancellation to stop handshake backoff") + } +} + +func TestWebSocketHandshakeRetriesShareOperationTimeout(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if attempts.Add(1) == 1 { + time.Sleep(600 * time.Millisecond) + http.Error(w, "temporarily unavailable", http.StatusServiceUnavailable) + return + } + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + if _, _, err := connection.ReadMessage(); err != nil { + return + } + time.Sleep(500 * time.Millisecond) + _ = connection.WriteJSON(map[string]any{"type": "state", "data": map[string]any{}}) + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 1, nil) + session.handshakeRetryDelays = []time.Duration{time.Millisecond, time.Millisecond} + defer session.Close() + started := time.Now() + if _, err := session.Call(t.Context(), "state", ""); err == nil { + t.Fatal("expected shared operation timeout to expire") + } + if elapsed := time.Since(started); elapsed > 1500*time.Millisecond { + t.Fatalf("expected handshake retries to share the operation timeout, took %s", elapsed) + } +} + +func TestRetryableWebSocketHandshakeStatuses(t *testing.T) { + retryable := []int{408, 429, 500, 502, 503, 504} + for _, statusCode := range retryable { + if !isRetryableWebSocketHandshakeStatus(statusCode) { + t.Errorf("expected status %d to be retryable", statusCode) + } + } + for _, statusCode := range []int{400, 401, 403, 404} { + if isRetryableWebSocketHandshakeStatus(statusCode) { + t.Errorf("expected status %d not to be retryable", statusCode) + } + } +} + +func TestRetryableWebSocketHandshakeErrors(t *testing.T) { + if !isRetryableWebSocketHandshakeError(&net.DNSError{Err: "temporary failure", IsTemporary: true}) { + t.Fatal("expected network error to be retryable") + } + if isRetryableWebSocketHandshakeError(errors.New("invalid proxy configuration")) { + t.Fatal("expected local configuration error not to be retryable") + } +} + +func TestRuntimeWebSocketURL(t *testing.T) { + tests := []struct { + baseURL string + want string + }{ + { + baseURL: "https://example.test/openenv?api-version=1", + want: "wss://example.test/openenv/ws?api-version=1", + }, + { + baseURL: "http://127.0.0.1:8080/openenv/", + want: "ws://127.0.0.1:8080/openenv/ws", + }, + } + for _, test := range tests { + got, err := RuntimeWebSocketURL(test.baseURL) + if err != nil { + t.Fatal(err) + } + if got != test.want { + t.Fatalf("RuntimeWebSocketURL(%q) = %q, want %q", test.baseURL, got, test.want) + } + } + if _, err := RuntimeWebSocketURL("ftp://example.test/openenv"); err == nil { + t.Fatal("expected unsupported scheme to fail") + } +} + +func TestWebSocketSessionSendsKeepalivePings(t *testing.T) { + pingReceived := make(chan struct{}, 1) + serverDone := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer close(serverDone) + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + connection.SetPingHandler(func(data string) error { + select { + case pingReceived <- struct{}{}: + default: + } + return connection.WriteControl( + websocket.PongMessage, + []byte(data), + time.Now().Add(time.Second), + ) + }) + for { + var request map[string]any + if err := connection.ReadJSON(&request); err != nil { + return + } + if err := connection.WriteJSON(map[string]any{ + "type": "state", + "data": map[string]any{}, + }); err != nil { + t.Errorf("write WebSocket response: %v", err) + return + } + } + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + session.keepAliveInterval = 10 * time.Millisecond + if _, err := session.Call(t.Context(), "state", ""); err != nil { + t.Fatal(err) + } + select { + case <-pingReceived: + case <-time.After(time.Second): + t.Fatal("expected WebSocket keepalive ping") + } + session.Close() + <-serverDone +} + +func TestWebSocketRequestUsesOpenEnvProtocol(t *testing.T) { + reset, err := webSocketRequest("reset", &callOptions{}) + if err != nil { + t.Fatal(err) + } + + step, err := webSocketRequest("step", &callOptions{action: `{"message":"hello"}`}) + if err != nil { + t.Fatal(err) + } + state, err := webSocketRequest("state", &callOptions{}) + if err != nil { + t.Fatal(err) + } + + assertJSONEqual(t, reset, `{"type":"reset","data":{}}`) + assertJSONEqual(t, step, `{"type":"step","data":{"message":"hello"}}`) + assertJSONEqual(t, state, `{"type":"state"}`) +} + +func TestParseWebSocketResponse(t *testing.T) { + result, terminal, err := parseWebSocketResponse( + "step", + []byte(`{"type":"observation","data":{"reward":1}}`), + ) + if err != nil || terminal || !strings.Contains(result, `"reward": 1`) { + t.Fatalf("unexpected successful response: result=%q terminal=%t err=%v", result, terminal, err) + } + + _, terminal, err = parseWebSocketResponse( + "step", + []byte(`{"type":"error","data":{"detail":"invalid action"}}`), + ) + if err == nil || terminal || !strings.Contains(err.Error(), "invalid action") { + t.Fatalf("unexpected OpenEnv error response: terminal=%t err=%v", terminal, err) + } + + for _, code := range []string{"CAPACITY_REACHED", "FACTORY_ERROR", "SESSION_ERROR"} { + _, terminal, err = parseWebSocketResponse( + "step", + fmt.Appendf(nil, `{"type":"error","data":{"code":%q,"detail":"failed"}}`, code), + ) + localError, ok := errors.AsType[*azdext.LocalError](err) + if err == nil || !terminal || !ok || + !strings.Contains(localError.Suggestion, "run invoke again") { + t.Fatalf("expected %s to be terminal with reinvoke guidance: terminal=%t err=%v", code, terminal, err) + } + } + + _, terminal, err = parseWebSocketResponse( + "step", + []byte(`{"type":"error","data":{"code":"VALIDATION_ERROR","detail":"invalid action"}}`), + ) + localError, ok := errors.AsType[*azdext.LocalError](err) + if err == nil || terminal || !ok || + !strings.Contains(localError.Suggestion, "payload and retry") { + t.Fatalf("expected validation error to be recoverable: terminal=%t err=%v", terminal, err) + } + + _, terminal, err = parseWebSocketResponse( + "state", + []byte(`{"type":"observation","data":{}}`), + ) + if err == nil || !terminal { + t.Fatalf("expected mismatched response type to be terminal, got terminal=%t err=%v", terminal, err) + } +} + +func TestWebSocketSessionErrorPreventsReconnect(t *testing.T) { + upgrades := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + upgrades++ + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer connection.Close() + if _, _, err := connection.ReadMessage(); err != nil { + t.Error(err) + return + } + if err := connection.WriteMessage( + websocket.TextMessage, + []byte(`{"type":"error","data":{"code":"SESSION_ERROR","detail":"session failed"}}`), + ); err != nil { + t.Error(err) + } + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + defer session.Close() + if _, err := session.Call(t.Context(), "state", ""); err == nil { + t.Fatal("expected the session error to fail the first call") + } + if _, err := session.Call(t.Context(), "state", ""); err == nil { + t.Fatal("expected the terminal session error to fail the second call") + } + if upgrades != 1 { + t.Fatalf("expected no reconnection after session error, got %d connections", upgrades) + } +} + +func TestWebSocketSessionCloseInterruptsCallAndPreventsReconnect(t *testing.T) { + requestReceived := make(chan struct{}) + upgrades := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + upgrades++ + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + + defer connection.Close() + if _, _, err := connection.ReadMessage(); err != nil { + t.Errorf("read WebSocket request: %v", err) + return + } + close(requestReceived) + _, _, _ = connection.ReadMessage() + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + callDone := make(chan error, 1) + go func() { + _, err := session.Call(t.Context(), "state", "") + callDone <- err + }() + <-requestReceived + session.Close() + if err := <-callDone; err == nil { + t.Fatal("expected closing the session to fail the active call") + } + if _, err := session.Call(t.Context(), "state", ""); err == nil { + t.Fatal("expected a closed session to reject subsequent calls") + } + if upgrades != 1 { + t.Fatalf("expected no reconnection after close, got %d connections", upgrades) + } +} + +func TestCanceledQueuedCallDoesNotReachWebSocket(t *testing.T) { + firstRequestReceived := make(chan struct{}) + releaseFirstRequest := make(chan struct{}) + requests := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + for { + var request map[string]any + if err := connection.ReadJSON(&request); err != nil { + if websocket.IsCloseError(err, websocket.CloseNormalClosure) { + return + } + t.Errorf("read WebSocket request: %v", err) + return + } + requests++ + if requests == 1 { + close(firstRequestReceived) + <-releaseFirstRequest + } + if err := connection.WriteJSON(map[string]any{ + "type": "state", + "data": map[string]any{"request": requests}, + }); err != nil { + t.Errorf("write WebSocket response: %v", err) + return + } + } + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 30, nil) + defer session.Close() + firstCallDone := make(chan error, 1) + go func() { + _, err := session.Call(t.Context(), "state", "") + firstCallDone <- err + }() + <-firstRequestReceived + + queuedContext, cancelQueued := context.WithCancel(t.Context()) + queuedCallDone := make(chan error, 1) + go func() { + _, err := session.CallAndDrain(queuedContext, "state", "") + queuedCallDone <- err + }() + cancelQueued() + close(releaseFirstRequest) + + if err := <-firstCallDone; err != nil { + t.Fatalf("first call failed: %v", err) + } + if err := <-queuedCallDone; err == nil { + t.Fatal("expected canceled queued call to fail") + } + if requests != 1 { + t.Fatalf("expected canceled queued call not to reach WebSocket, got %d requests", requests) + } +} + +func TestCallAndDrainHasNoImplicitTimeoutWhenDisabled(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + if _, _, err := connection.ReadMessage(); err != nil { + t.Errorf("read WebSocket request: %v", err) + return + } + time.Sleep(25 * time.Millisecond) + if err := connection.WriteJSON(map[string]any{"type": "state", "data": map[string]any{}}); err != nil { + t.Errorf("write WebSocket response: %v", err) + } + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 0, nil) + defer session.Close() + if _, err := session.CallAndDrain(t.Context(), "state", ""); err != nil { + t.Fatalf("expected timeout-disabled call to wait for its response: %v", err) + } +} + +func TestCallAndDrainBoundsResponseDrainAfterCancellation(t *testing.T) { + requestReceived := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + connection, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade WebSocket: %v", err) + return + } + defer connection.Close() + if _, _, err := connection.ReadMessage(); err != nil { + t.Errorf("read WebSocket request: %v", err) + return + } + close(requestReceived) + _, _, _ = connection.ReadMessage() + })) + defer server.Close() + + session := NewWebSocketRuntimeSession(server.URL, 0, nil) + session.drainTimeout = 10 * time.Millisecond + defer session.Close() + ctx, cancel := context.WithCancel(t.Context()) + callDone := make(chan error, 1) + go func() { + _, err := session.CallAndDrain(ctx, "state", "") + callDone <- err + }() + <-requestReceived + cancel() + select { + case err := <-callDone: + if err == nil { + t.Fatal("expected response drain to end after cancellation") + } + case <-time.After(time.Second): + t.Fatal("expected cancellation to bound the response drain") + } +} + +func assertWebSocketRequest( + t *testing.T, + request map[string]any, + requestType string, + data map[string]any, +) { + t.Helper() + if request["type"] != requestType { + t.Fatalf("request type = %#v, want %q", request["type"], requestType) + } + if data == nil { + if _, ok := request["data"]; ok { + t.Fatalf("did not expect data in %#v", request) + } + return + } + actual, ok := request["data"].(map[string]any) + if !ok || !mapsEqual(actual, data) { + t.Fatalf("request data = %#v, want %#v", request["data"], data) + } +} + +func assertJSONEqual(t *testing.T, actual []byte, expected string) { + t.Helper() + var actualValue any + var expectedValue any + if err := json.Unmarshal(actual, &actualValue); err != nil { + t.Fatal(err) + } + if err := json.Unmarshal([]byte(expected), &expectedValue); err != nil { + t.Fatal(err) + } + actualJSON, _ := json.Marshal(actualValue) + expectedJSON, _ := json.Marshal(expectedValue) + if !bytes.Equal(actualJSON, expectedJSON) { + t.Fatalf("JSON = %s, want %s", actual, expected) + } +} + +func mapsEqual(left map[string]any, right map[string]any) bool { + leftJSON, _ := json.Marshal(left) + rightJSON, _ := json.Marshal(right) + return bytes.Equal(leftJSON, rightJSON) +} diff --git a/cli/azd/extensions/azure.ai.rle/version.txt b/cli/azd/extensions/azure.ai.rle/version.txt index 6bcee0a14ab..34dab3faf78 100644 --- a/cli/azd/extensions/azure.ai.rle/version.txt +++ b/cli/azd/extensions/azure.ai.rle/version.txt @@ -1 +1 @@ -0.4.0-preview \ No newline at end of file +0.4.1-preview \ No newline at end of file