diff --git a/internal/app/install.go b/internal/app/install.go index 8b534758..48b5e0b6 100644 --- a/internal/app/install.go +++ b/internal/app/install.go @@ -236,15 +236,12 @@ func (u *UseCases) probeInstallProtocols(ctx context.Context, options InstallAge ordered = append(ordered, protocolID) } sort.Strings(ordered) - for _, protocolID := range ordered { - if u.provider == nil { - return nil, oneerrors.New(oneerrors.InternalError, "Provider probing is not configured", oneerrors.WithStatus(501)) - } - verdict, err := u.provider.Probe(ctx, protocolID, "custom", options.APIKey, options.Model, target.BaseFor(protocolID)) - if err != nil { - return nil, err - } - probes[protocolID] = verdict + if u.provider == nil { + return nil, oneerrors.New(oneerrors.InternalError, "Provider probing is not configured", oneerrors.WithStatus(501)) + } + probes, err = u.probeProtocols(ctx, ordered, options.APIKey, options.Model, target.BaseFor) + if err != nil { + return nil, err } u.sharpenInstallModelDiagnosis(ctx, probes, options) return probes, nil diff --git a/internal/app/provider.go b/internal/app/provider.go index f86b4a11..274e734b 100644 --- a/internal/app/provider.go +++ b/internal/app/provider.go @@ -5,6 +5,7 @@ import ( "fmt" "sort" "strings" + "sync" "github.com/MaimoryLab/OneAgent/internal/catalog" "github.com/MaimoryLab/OneAgent/internal/desktopapp" @@ -57,13 +58,9 @@ func (u *UseCases) ProbeProvider(ctx context.Context, options ProviderProbeOptio if err != nil { return ProviderProbeResult{}, err } - results := make(map[string]provider.ProbeResult, len(protocols)) - for _, protocolID := range protocols { - result, probeErr := u.provider.Probe(ctx, protocolID, "custom", apiKey, model, target.BaseFor(protocolID)) - if probeErr != nil { - return ProviderProbeResult{}, probeErr - } - results[protocolID] = result + results, err := u.probeProtocols(ctx, protocols, apiKey, model, target.BaseFor) + if err != nil { + return ProviderProbeResult{}, err } primary := results[protocols[0]] allOK := true @@ -79,6 +76,34 @@ func (u *UseCases) ProbeProvider(ctx context.Context, options ProviderProbeOptio return ProviderProbeResult{Primary: primary, Protocols: results}, nil } +func (u *UseCases) probeProtocols(ctx context.Context, protocols []string, apiKey, model string, baseFor func(string) string) (map[string]provider.ProbeResult, error) { + results := make(map[string]provider.ProbeResult, len(protocols)) + errorsByProtocol := make(map[string]error) + var mu sync.Mutex + var group sync.WaitGroup + for _, protocolID := range protocols { + group.Add(1) + go func(protocolID string) { + defer group.Done() + result, err := u.provider.Probe(ctx, protocolID, "custom", apiKey, model, baseFor(protocolID)) + mu.Lock() + defer mu.Unlock() + if err != nil { + errorsByProtocol[protocolID] = err + return + } + results[protocolID] = result + }(protocolID) + } + group.Wait() + for _, protocolID := range protocols { + if err := errorsByProtocol[protocolID]; err != nil { + return results, err + } + } + return results, nil +} + func (u *UseCases) ListProviderModels(ctx context.Context, providerID, apiKey, customBase string) (provider.ModelsResult, error) { if err := ctx.Err(); err != nil { return provider.ModelsResult{}, oneerrors.New(oneerrors.Timeout, "Request was cancelled", oneerrors.WithRetryable(true), oneerrors.WithCause(err)) diff --git a/internal/app/provider_test.go b/internal/app/provider_test.go index 1800bdb0..2cb3b581 100644 --- a/internal/app/provider_test.go +++ b/internal/app/provider_test.go @@ -9,6 +9,7 @@ import ( "reflect" "sort" "strings" + "sync" "testing" oneerrors "github.com/MaimoryLab/OneAgent/internal/errors" @@ -38,7 +39,10 @@ func providerUseCases(t *testing.T, doer provider.HTTPDoer) *UseCases { func TestProbeProviderAggregatesAgentProtocols(t *testing.T) { seen := make([]string, 0) + var seenMu sync.Mutex core := providerUseCases(t, appProviderDoer(func(request *http.Request) (*http.Response, error) { + seenMu.Lock() + defer seenMu.Unlock() seen = append(seen, request.URL.Path) return appProviderResponse(http.StatusNoContent, ""), nil })) diff --git a/internal/app/status.go b/internal/app/status.go index 00e2eabf..af88dfbe 100644 --- a/internal/app/status.go +++ b/internal/app/status.go @@ -27,6 +27,8 @@ import ( type CommandLookup func(string) (string, bool) +const versionProbeConcurrency = 3 + type StatusOptions struct { Home string Platform platform.Info @@ -280,6 +282,7 @@ func (u *UseCases) GetStatus(ctx context.Context) (StatusResponse, error) { statuses := make(map[string]AgentStatus, len(manifest.Agents)) bindings := u.profiles.ListAgentBindings() latestVersions := u.latestAgentVersions(ctx, manifest, agentLookup) + installedVersions := u.installedVersions(ctx, manifest, agentLookup) for _, id := range catalog.AgentIDs(manifest) { agent := manifest.Agents[id] configPath := configPath(options.Home, options.Platform.OS, agent) @@ -287,9 +290,8 @@ func (u *UseCases) GetStatus(ctx context.Context) (StatusResponse, error) { paths[id+"_config"] = configPath } installed := false - executable := "" if agent.Command != "" { - executable, installed = agentLookup(agent.Command) + _, installed = agentLookup(agent.Command) } canInstall := false if agent.Package != nil { @@ -324,7 +326,7 @@ func (u *UseCases) GetStatus(ctx context.Context) (StatusResponse, error) { } var installedVersion *string if installed && agent.ConfigMode == "auto" { - installedVersion = u.installedVersion(ctx, executable, agent.VersionArgs) + installedVersion = installedVersions[id] } statuses[id] = AgentStatus{ Installed: installed, @@ -370,6 +372,50 @@ func (u *UseCases) GetStatus(ctx context.Context) (StatusResponse, error) { }, nil } +func (u *UseCases) installedVersions(ctx context.Context, manifest catalog.Manifest, lookup func(string) (string, bool)) map[string]*string { + queries := make([]struct { + id, executable string + args []string + }, 0, len(manifest.Agents)) + for id, agent := range manifest.Agents { + if agent.ConfigMode != "auto" || agent.Command == "" { + continue + } + executable, installed := lookup(agent.Command) + if installed { + queries = append(queries, struct { + id, executable string + args []string + }{id, executable, agent.VersionArgs}) + } + } + if len(queries) == 0 { + return nil + } + versions := make(map[string]*string, len(queries)) + var mu sync.Mutex + var group sync.WaitGroup + tokens := make(chan struct{}, versionProbeConcurrency) + for _, query := range queries { + group.Add(1) + go func(query struct { + id, executable string + args []string + }) { + defer group.Done() + tokens <- struct{}{} + defer func() { <-tokens }() + if version := u.installedVersion(ctx, query.executable, query.args); version != nil { + mu.Lock() + versions[query.id] = version + mu.Unlock() + } + }(query) + } + group.Wait() + return versions +} + var versionPattern = regexp.MustCompile(`(^|[^\d])(\d+\.\d+\.\d+(?:[-+][0-9A-Za-z.-]+)?)`) // installedVersion runs the Agent's version command and takes the first diff --git a/internal/app/status_test.go b/internal/app/status_test.go index 390e495b..d4ceada3 100644 --- a/internal/app/status_test.go +++ b/internal/app/status_test.go @@ -3,13 +3,18 @@ package app import ( "context" "encoding/json" + "fmt" "os" "path/filepath" "reflect" "strings" + "sync/atomic" "testing" + "time" + "github.com/MaimoryLab/OneAgent/internal/catalog" "github.com/MaimoryLab/OneAgent/internal/platform" + "github.com/MaimoryLab/OneAgent/internal/process" "github.com/MaimoryLab/OneAgent/internal/provider" ) @@ -301,6 +306,35 @@ func TestStatusReportsInstalledVersionFromVersionCommand(t *testing.T) { } } +func TestInstalledVersionsProbeConcurrentlyWithBoundedFanout(t *testing.T) { + runner := &versionProbeRunner{} + core := NewUseCases(StatusOptions{Runner: runner, Platform: platform.For("linux", "amd64")}) + manifest := catalog.Manifest{Agents: map[string]catalog.Agent{}} + for index := 0; index < versionProbeConcurrency+1; index++ { + manifest.Agents[fmt.Sprintf("agent-%d", index)] = catalog.Agent{Command: fmt.Sprintf("cmd-%d", index), ConfigMode: "auto"} + } + versions := core.installedVersions(context.Background(), manifest, runner.LookPath) + if len(versions) != versionProbeConcurrency+1 || runner.peak.Load() != versionProbeConcurrency { + t.Fatalf("versions=%d peak=%d, want %d and %d", len(versions), runner.peak.Load(), versionProbeConcurrency+1, versionProbeConcurrency) + } +} + +type versionProbeRunner struct { + inFlight atomic.Int32 + peak atomic.Int32 +} + +func (r *versionProbeRunner) LookPath(command string) (string, bool) { return "/fake/" + command, true } + +func (r *versionProbeRunner) Run(_ context.Context, argv []string, _ map[string]string, _ time.Duration) (process.Result, error) { + current := r.inFlight.Add(1) + for current > r.peak.Load() && !r.peak.CompareAndSwap(r.peak.Load(), current) { + } + time.Sleep(10 * time.Millisecond) + r.inFlight.Add(-1) + return process.Result{Stdout: "tool 1.0.0"}, nil +} + func TestStatusMatchesEmptyLinuxARM64Fixture(t *testing.T) { home := t.TempDir() core := NewUseCases(StatusOptions{ diff --git a/internal/binding/services_test.go b/internal/binding/services_test.go index 549e1f5d..cf53b235 100644 --- a/internal/binding/services_test.go +++ b/internal/binding/services_test.go @@ -10,6 +10,7 @@ import ( "reflect" "sort" "strings" + "sync" "testing" "github.com/MaimoryLab/OneAgent/internal/app" @@ -331,7 +332,10 @@ func TestInstallResultBindingPreservesFieldPresence(t *testing.T) { func TestProviderServiceAggregatesSelectedAgentProtocols(t *testing.T) { seen := make([]string, 0) + var seenMu sync.Mutex client := provider.NewClient(providerFakeDoer(func(request *http.Request) (*http.Response, error) { + seenMu.Lock() + defer seenMu.Unlock() seen = append(seen, request.URL.Path) return providerResponse(http.StatusNoContent, ""), nil }))