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
15 changes: 6 additions & 9 deletions internal/app/install.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
39 changes: 32 additions & 7 deletions internal/app/provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"fmt"
"sort"
"strings"
"sync"

"github.com/MaimoryLab/OneAgent/internal/catalog"
"github.com/MaimoryLab/OneAgent/internal/desktopapp"
Expand Down Expand Up @@ -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
Expand All @@ -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))
Expand Down
4 changes: 4 additions & 0 deletions internal/app/provider_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
"reflect"
"sort"
"strings"
"sync"
"testing"

oneerrors "github.com/MaimoryLab/OneAgent/internal/errors"
Expand Down Expand Up @@ -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
}))
Expand Down
52 changes: 49 additions & 3 deletions internal/app/status.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,8 @@ import (

type CommandLookup func(string) (string, bool)

const versionProbeConcurrency = 3

type StatusOptions struct {
Home string
Platform platform.Info
Expand Down Expand Up @@ -280,16 +282,16 @@ 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)
if configPath != "" {
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 {
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
34 changes: 34 additions & 0 deletions internal/app/status_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down Expand Up @@ -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{
Expand Down
4 changes: 4 additions & 0 deletions internal/binding/services_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"reflect"
"sort"
"strings"
"sync"
"testing"

"github.com/MaimoryLab/OneAgent/internal/app"
Expand Down Expand Up @@ -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
}))
Expand Down