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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 19 additions & 1 deletion cli/azd/extensions/azure.ai.agents/internal/cmd/run.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ import (
const (
agentInspectorExtensionID = "azure.ai.inspector"
agentInspectorReadyPollPeriod = 250 * time.Millisecond
windowsControlCExitCode = 0xC000013A
// defaultInspectorUIPort mirrors the default UI port of the
// azure.ai.inspector extension. The inspector extension remains the source
// of truth for the actual default: when --inspector-port is unset we do not
Expand Down Expand Up @@ -364,11 +365,12 @@ func runRun(ctx context.Context, flags *runFlags, noPrompt bool) error {

err = proc.Wait()
close(done)
wasCanceled := ctx.Err() != nil || isAgentProcessCanceled(err)
cancel()
<-nextDone

// Suppress the noisy "signal: interrupt" error on Ctrl+C
if ctx.Err() != nil {
if wasCanceled {
fmt.Println("Agent stopped.")
return nil
}
Expand All @@ -379,6 +381,22 @@ func runRun(ctx context.Context, flags *runFlags, noPrompt bool) error {
return nil
}

func isAgentProcessCanceled(err error) bool {
exitErr, ok := errors.AsType[*exec.ExitError](err)
if !ok {
return false
}

if waitStatus, ok := exitErr.Sys().(syscall.WaitStatus); ok {
switch waitStatus.Signal() {
case syscall.SIGINT, syscall.SIGTERM:
return true
}
}

return uint32(exitErr.ExitCode()) == windowsControlCExitCode //nolint:gosec // preserve the Windows exit code bit pattern
}

func handleInspectorAutoLaunch(
ctx context.Context,
workflow azdext.WorkflowServiceClient,
Expand Down
103 changes: 103 additions & 0 deletions cli/azd/extensions/azure.ai.agents/internal/cmd/run_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -529,6 +529,109 @@ func TestRunRun_PortCollisionDoesNotClearStoredSession(t *testing.T) {
}
}

func TestRunRun_ReturnsAgentProcessExitError(t *testing.T) {
err := runRunWithHelperProcess(t, "exit", "17")
if err == nil || !strings.Contains(err.Error(), "agent exited: exit status 17") {
t.Fatalf("runRun() error = %v, want agent exit status 17", err)
}
}

func TestRunRun_TreatsInterruptExitAsCancellation(t *testing.T) {
if err := runRunWithHelperProcess(t, "interrupt", ""); err != nil {
t.Fatalf("runRun() error = %v, want nil for interrupt exit", err)
}
}

func runRunWithHelperProcess(t *testing.T, mode string, exitCode string) error {
t.Helper()

projectDir := t.TempDir()
projectServer := &helpersProjectServer{
project: &azdext.ProjectConfig{
Name: "test-project",
Path: projectDir,
Services: map[string]*azdext.ServiceConfig{
"agent": {
Name: "agent",
Host: AiAgentHost,
RelativePath: ".",
},
},
},
}

grpcServer := grpc.NewServer()
azdext.RegisterProjectServiceServer(grpcServer, projectServer)
azdext.RegisterUserConfigServiceServer(grpcServer, newInvokeUserConfigServer())

listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen: %v", err)
}
go func() { _ = grpcServer.Serve(listener) }()
t.Cleanup(func() {
grpcServer.Stop()
_ = listener.Close()
})
t.Setenv("AZD_SERVER", listener.Addr().String())
t.Setenv("AZD_AGENT_RUN_TEST_HELPER_MODE", mode)
t.Setenv("AZD_AGENT_RUN_TEST_EXIT_CODE", exitCode)

agentListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("reserve agent port: %v", err)
}
agentPort := agentListener.Addr().(*net.TCPAddr).Port
if err := agentListener.Close(); err != nil {
t.Fatalf("release agent port: %v", err)
}

startCommand := fmt.Sprintf(`"%s" -test.run=^TestRunRunHelperProcess$`, os.Args[0])
return runRun(t.Context(), &runFlags{
name: "agent",
port: agentPort,
startCommand: startCommand,
noClient: true,
}, true)
}

func TestRunRunHelperProcess(t *testing.T) {
mode := os.Getenv("AZD_AGENT_RUN_TEST_HELPER_MODE")
if mode == "" {
return
}

if mode == "interrupt" {
if runtime.GOOS == "windows" {
exitCode := uint32(windowsControlCExitCode)
os.Exit(int(exitCode)) //nolint:gosec // preserve the Windows exit code bit pattern
}

time.AfterFunc(5*time.Second, func() {
os.Exit(99)
})
proc, err := os.FindProcess(os.Getpid())
if err != nil {
t.Fatalf("find helper process: %v", err)
}
if err := proc.Signal(os.Interrupt); err != nil {
t.Fatalf("interrupt helper process: %v", err)
}
select {}
}

exitCodeValue := os.Getenv("AZD_AGENT_RUN_TEST_EXIT_CODE")
if exitCodeValue == "" {
t.Fatalf("missing helper exit code for mode %q", mode)
}

exitCode, err := strconv.Atoi(exitCodeValue)
if err != nil {
t.Fatalf("parse helper exit code: %v", err)
}
os.Exit(exitCode)
}

func TestWarnInspectorPortIssues(t *testing.T) {
t.Parallel()

Expand Down
Loading