diff --git a/cmd/mcp-sim/main.go b/cmd/mcp-sim/main.go index c436aed..15a2d82 100644 --- a/cmd/mcp-sim/main.go +++ b/cmd/mcp-sim/main.go @@ -9,7 +9,9 @@ import ( "os" "os/signal" "sort" + "strings" "syscall" + "time" "github.com/espetro/mcp-sim/controllers/agentdevice" "github.com/espetro/mcp-sim/internal/config" @@ -182,7 +184,16 @@ func serveImpl(prog, listenAddr, configPath string) error { ctx := applog.WithContext(context.Background(), logger) registry := core.NewRegistry(logger) - lifecycle := core.NewLifecycle(registry) + + timeout, err := parseSessionIdleTimeout(cfg.Server.SessionIdleTimeout) + if err != nil { + return fmt.Errorf("parsing session_idle_timeout: %w", err) + } + sessionManager := core.NewManager(timeout, logger) + lifecycle := core.NewLifecycle(registry, sessionManager) + sessionManager.SetStopper(func(ctx context.Context, platform, target, owner string) error { + return lifecycle.StopDevice(ctx, platform, target, owner) + }) if cfg.Platforms.IOS.Enabled { iosPlatform, err := ios.New(ctx, cfg.Platforms.IOS) @@ -214,7 +225,7 @@ func serveImpl(prog, listenAddr, configPath string) error { } } - mcpServer := mcp.NewServer(registry, lifecycle, logger) + mcpServer := mcp.NewServer(registry, lifecycle, sessionManager, logger) mux := http.NewServeMux() mux.Handle("/mcp", mcpServer.StreamableHTTPHandler()) @@ -233,7 +244,9 @@ func serveImpl(prog, listenAddr, configPath string) error { ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) defer stop() + ctx = applog.WithContext(ctx, logger) + go sessionManager.Start(ctx) go func() { <-ctx.Done() logger.Info("shutdown signal received") @@ -257,7 +270,16 @@ func mcpImpl(prog, configPath string) error { ctx := applog.WithContext(context.Background(), logger) registry := core.NewRegistry(logger) - lifecycle := core.NewLifecycle(registry) + + timeout, err := parseSessionIdleTimeout(cfg.Server.SessionIdleTimeout) + if err != nil { + return fmt.Errorf("parsing session_idle_timeout: %w", err) + } + sessionManager := core.NewManager(timeout, logger) + lifecycle := core.NewLifecycle(registry, sessionManager) + sessionManager.SetStopper(func(ctx context.Context, platform, target, owner string) error { + return lifecycle.StopDevice(ctx, platform, target, owner) + }) if cfg.Platforms.IOS.Enabled { iosPlatform, _ := ios.New(ctx, cfg.Platforms.IOS) @@ -277,7 +299,9 @@ func mcpImpl(prog, configPath string) error { } } - mcpServer := mcp.NewServer(registry, lifecycle, logger) + go sessionManager.Start(ctx) + + mcpServer := mcp.NewServer(registry, lifecycle, sessionManager, logger) return mcpServer.Run(ctx, &sdkmcp.StdioTransport{}) } @@ -301,3 +325,14 @@ func controllerNames(r *core.Registry) []string { sort.Strings(names) return names } + +func parseSessionIdleTimeout(v string) (time.Duration, error) { + v = strings.TrimSpace(strings.ToLower(v)) + if v == "" { + return 30 * time.Minute, nil + } + if v == "0" || v == "off" { + return 0, nil + } + return time.ParseDuration(v) +} diff --git a/internal/config/config.go b/internal/config/config.go index cec7f8b..602a139 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -17,9 +17,10 @@ type Config struct { // ServerConfig configures the HTTP server. type ServerConfig struct { - Listen string `yaml:"listen"` // MCPSIM_LISTEN - LogLevel string `yaml:"log_level"` // MCPSIM_LOG_LEVEL - LogFormat string `yaml:"log_format"` // MCPSIM_LOG_FORMAT + Listen string `yaml:"listen"` // MCPSIM_LISTEN + LogLevel string `yaml:"log_level"` // MCPSIM_LOG_LEVEL + LogFormat string `yaml:"log_format"` // MCPSIM_LOG_FORMAT + SessionIdleTimeout string `yaml:"session_idle_timeout"` // MCPSIM_SESSION_IDLE_TIMEOUT } // PlatformsConfig holds per-platform configuration. @@ -57,9 +58,10 @@ type AgentDeviceConfig struct { func defaultConfig() Config { return Config{ Server: ServerConfig{ - Listen: ":9090", - LogLevel: "info", - LogFormat: "text", + Listen: ":9090", + LogLevel: "info", + LogFormat: "text", + SessionIdleTimeout: "30m", }, Platforms: PlatformsConfig{ IOS: IOSConfig{ @@ -103,6 +105,9 @@ func Load() (Config, error) { if v := os.Getenv("MCPSIM_LOG_FORMAT"); v != "" { cfg.Server.LogFormat = v } + if v := os.Getenv("MCPSIM_SESSION_IDLE_TIMEOUT"); v != "" { + cfg.Server.SessionIdleTimeout = v + } if v := os.Getenv("MCPSIM_IOS_ENABLED"); v != "" { cfg.Platforms.IOS.Enabled, _ = strconv.ParseBool(v) } diff --git a/internal/core/lifecycle.go b/internal/core/lifecycle.go index 043b083..26c8689 100644 --- a/internal/core/lifecycle.go +++ b/internal/core/lifecycle.go @@ -10,20 +10,30 @@ import ( // Lifecycle orchestrates boot/shutdown of devices with state management. type Lifecycle struct { registry *Registry + sessions *Manager } // NewLifecycle creates a new lifecycle orchestrator. -func NewLifecycle(registry *Registry) *Lifecycle { - return &Lifecycle{registry: registry} +func NewLifecycle(registry *Registry, sessions *Manager) *Lifecycle { + return &Lifecycle{registry: registry, sessions: sessions} } // BootDevice boots a device, handling already-running and not-found errors. -func (l *Lifecycle) BootDevice(ctx context.Context, platformName, target string, opts contract.StartOpts) (contract.Device, error) { +func (l *Lifecycle) BootDevice(ctx context.Context, platformName, target, sessionID string, opts contract.StartOpts) (dev contract.Device, err error) { p, ok := l.registry.PlatformByName(platformName) if !ok { return contract.Device{}, &ToolError{Code: contract.ErrUnsupportedPlatform, Msg: "platform not found: " + platformName} } + if err := l.sessions.Reserve(ctx, platformName, target, sessionID); err != nil { + return contract.Device{}, err + } + defer func() { + if err != nil { + l.sessions.Release(platformName, target) + } + }() + state, err := p.State(ctx, target) if err != nil { return contract.Device{}, err @@ -33,7 +43,7 @@ func (l *Lifecycle) BootDevice(ctx context.Context, platformName, target string, return contract.Device{}, &ToolError{Code: contract.ErrAlreadyRunning, Msg: "device already running: " + target} } - dev, err := p.Start(ctx, target, opts) + dev, err = p.Start(ctx, target, opts) if err != nil { return contract.Device{}, err } @@ -48,16 +58,21 @@ func (l *Lifecycle) BootDevice(ctx context.Context, platformName, target string, return dev, &ToolError{Code: contract.ErrTimeout, Msg: "device did not become ready: " + target} } + l.sessions.RecordActivity(platformName, target) return dev, nil } // StopDevice stops a device. -func (l *Lifecycle) StopDevice(ctx context.Context, platformName, target string) error { +func (l *Lifecycle) StopDevice(ctx context.Context, platformName, target, sessionID string) error { p, ok := l.registry.PlatformByName(platformName) if !ok { return &ToolError{Code: contract.ErrUnsupportedPlatform, Msg: "platform not found: " + platformName} } + if err := l.sessions.CheckAccess(platformName, target, sessionID); err != nil { + return err + } + state, err := p.State(ctx, target) if err != nil { return err @@ -67,16 +82,26 @@ func (l *Lifecycle) StopDevice(ctx context.Context, platformName, target string) return &ToolError{Code: contract.ErrNotRunning, Msg: "device not running: " + target} } - return p.Stop(ctx, target) + if err := p.Stop(ctx, target); err != nil { + return err + } + + l.sessions.Release(platformName, target) + return nil } // WipeDevice wipes a device. -func (l *Lifecycle) WipeDevice(ctx context.Context, platformName, target string) error { +func (l *Lifecycle) WipeDevice(ctx context.Context, platformName, target, sessionID string) error { p, ok := l.registry.PlatformByName(platformName) if !ok { return &ToolError{Code: contract.ErrUnsupportedPlatform, Msg: "platform not found: " + platformName} } + if err := l.sessions.CheckAccess(platformName, target, sessionID); err != nil { + return err + } + defer l.sessions.Release(platformName, target) + state, err := p.State(ctx, target) if err != nil { return err @@ -92,10 +117,15 @@ func (l *Lifecycle) WipeDevice(ctx context.Context, platformName, target string) } // OpenURL opens a deep link on a device. -func (l *Lifecycle) OpenURL(ctx context.Context, platformName, target, url string) error { +func (l *Lifecycle) OpenURL(ctx context.Context, platformName, target, sessionID, url string) error { p, ok := l.registry.PlatformByName(platformName) if !ok { return &ToolError{Code: contract.ErrUnsupportedPlatform, Msg: "platform not found: " + platformName} } + + if err := l.sessions.CheckAccess(platformName, target, sessionID); err != nil { + return err + } + l.sessions.RecordActivity(platformName, target) return p.OpenURL(ctx, target, url) } diff --git a/internal/core/lifecycle_test.go b/internal/core/lifecycle_test.go new file mode 100644 index 0000000..fd9a12e --- /dev/null +++ b/internal/core/lifecycle_test.go @@ -0,0 +1,144 @@ +package core + +import ( + "context" + "errors" + "io" + "log/slog" + "sync" + "testing" + "time" + + "github.com/espetro/mcp-sim/pkg/contract" +) + +type fakePlatform struct { + mu sync.Mutex + name string + state contract.DeviceState + started bool + stopped bool +} + +func (p *fakePlatform) Name() string { return p.name } + +func (p *fakePlatform) List(ctx context.Context) ([]contract.Device, error) { + return []contract.Device{{ID: "dev1", Name: "dev1", Platform: p.name, State: p.state}}, nil +} + +func (p *fakePlatform) Start(ctx context.Context, target string, opts contract.StartOpts) (contract.Device, error) { + p.mu.Lock() + defer p.mu.Unlock() + p.state = contract.DeviceStateRunning + p.started = true + return contract.Device{ID: target, Name: target, Platform: p.name, State: p.state}, nil +} + +func (p *fakePlatform) Stop(ctx context.Context, target string) error { + p.mu.Lock() + defer p.mu.Unlock() + p.state = contract.DeviceStateStopped + p.stopped = true + return nil +} + +func (p *fakePlatform) State(ctx context.Context, target string) (contract.DeviceState, error) { + p.mu.Lock() + defer p.mu.Unlock() + return p.state, nil +} + +func (p *fakePlatform) AwaitReady(ctx context.Context, target string, timeout time.Duration) error { + return nil +} + +func (p *fakePlatform) Wipe(ctx context.Context, target string) error { + p.mu.Lock() + defer p.mu.Unlock() + p.state = contract.DeviceStateStopped + return nil +} + +func (p *fakePlatform) OpenURL(ctx context.Context, target, url string) error { + return nil +} + +func newTestLifecycle(p contract.Platform) (*Lifecycle, *Manager) { + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + reg := NewRegistry(logger) + reg.RegisterPlatform(p) + sessions := NewManager(time.Hour, logger) + return NewLifecycle(reg, sessions), sessions +} + +func TestBootDeviceReservesAndSecondSessionDenied(t *testing.T) { + fp := &fakePlatform{name: "ios", state: contract.DeviceStateStopped} + lc, sessions := newTestLifecycle(fp) + ctx := context.Background() + + if _, err := lc.BootDevice(ctx, "ios", "dev1", "session-a", contract.StartOpts{}); err != nil { + t.Fatalf("BootDevice failed: %v", err) + } + + owner, ok := sessions.Owner("ios", "dev1") + if !ok || owner != "session-a" { + t.Fatalf("Owner = (%q, %v), want (session-a, true)", owner, ok) + } + + _, err := lc.BootDevice(ctx, "ios", "dev1", "session-b", contract.StartOpts{}) + var te *ToolError + if !errors.As(err, &te) || te.Code != contract.ErrDeviceReserved { + t.Fatalf("BootDevice by other session returned %v, want device_reserved", err) + } +} + +func TestBootDeviceIdempotentSameSession(t *testing.T) { + fp := &fakePlatform{name: "ios", state: contract.DeviceStateStopped} + lc, _ := newTestLifecycle(fp) + ctx := context.Background() + + if _, err := lc.BootDevice(ctx, "ios", "dev1", "session-a", contract.StartOpts{}); err != nil { + t.Fatalf("first BootDevice failed: %v", err) + } + // Second boot by the same session is allowed through the reservation gate. + // Because the fake platform reports running, it returns already_running. + _, err := lc.BootDevice(ctx, "ios", "dev1", "session-a", contract.StartOpts{}) + var te *ToolError + if !errors.As(err, &te) || te.Code != contract.ErrAlreadyRunning { + t.Fatalf("second BootDevice returned %v, want already_running", err) + } +} + +func TestStopDeviceReleasesReservation(t *testing.T) { + fp := &fakePlatform{name: "ios", state: contract.DeviceStateRunning} + lc, sessions := newTestLifecycle(fp) + ctx := context.Background() + + if err := lc.StopDevice(ctx, "ios", "dev1", "session-a"); err != nil { + t.Fatalf("StopDevice failed: %v", err) + } + + if _, ok := sessions.Owner("ios", "dev1"); ok { + t.Fatal("reservation should be released after StopDevice") + } + + // After release, another session can claim the device. + if _, err := lc.BootDevice(ctx, "ios", "dev1", "session-b", contract.StartOpts{}); err != nil { + t.Fatalf("BootDevice by another session after stop failed: %v", err) + } +} + +func TestStopDeviceDeniedForOtherSession(t *testing.T) { + fp := &fakePlatform{name: "ios", state: contract.DeviceStateStopped} + lc, _ := newTestLifecycle(fp) + ctx := context.Background() + + // Boot reserves the device to session-a. + if _, err := lc.BootDevice(ctx, "ios", "dev1", "session-a", contract.StartOpts{}); err != nil { + t.Fatalf("BootDevice failed: %v", err) + } + + if err := lc.StopDevice(ctx, "ios", "dev1", "session-b"); err == nil { + t.Fatal("StopDevice by other session should fail") + } +} diff --git a/internal/core/session.go b/internal/core/session.go new file mode 100644 index 0000000..c544b6f --- /dev/null +++ b/internal/core/session.go @@ -0,0 +1,174 @@ +package core + +import ( + "context" + "fmt" + "log/slog" + "sync" + "time" + + "github.com/espetro/mcp-sim/pkg/contract" +) + +// sessionKey identifies a specific device reservation. +type sessionKey struct { + platform string + target string +} + +// Manager tracks per-session device reservations and idle timeouts. +type Manager struct { + mu sync.Mutex + reservations map[sessionKey]string + activity map[sessionKey]time.Time + timeout time.Duration + logger *slog.Logger + stopper func(context.Context, string, string, string) error + now func() time.Time +} + +// NewManager creates a reservation manager with the given idle timeout. +// A timeout of 0 disables the idle sweeper. +func NewManager(timeout time.Duration, logger *slog.Logger) *Manager { + return &Manager{ + reservations: make(map[sessionKey]string), + activity: make(map[sessionKey]time.Time), + timeout: timeout, + logger: logger, + now: time.Now, + } +} + +// SetStopper injects the callback used by the idle sweeper to stop a device. +// The callback receives the owning session so it can pass access checks. +func (m *Manager) SetStopper(stopper func(context.Context, string, string, string) error) { + m.mu.Lock() + defer m.mu.Unlock() + m.stopper = stopper +} + +// Reserve claims a device for a session. Idempotent if the same session +// already owns the device. +func (m *Manager) Reserve(ctx context.Context, platform, target, sessionID string) error { + _ = ctx + key := sessionKey{platform: platform, target: target} + m.mu.Lock() + defer m.mu.Unlock() + + owner, ok := m.reservations[key] + if ok && owner != sessionID { + return &ToolError{ + Code: contract.ErrDeviceReserved, + Msg: fmt.Sprintf("device reserved by session %s", owner), + } + } + + m.reservations[key] = sessionID + m.activity[key] = m.now() + return nil +} + +// CheckAccess returns an error if the device is reserved by another session. +func (m *Manager) CheckAccess(platform, target, sessionID string) error { + key := sessionKey{platform: platform, target: target} + m.mu.Lock() + defer m.mu.Unlock() + + owner, ok := m.reservations[key] + if ok && owner != sessionID { + return &ToolError{ + Code: contract.ErrDeviceReserved, + Msg: fmt.Sprintf("device reserved by session %s", owner), + } + } + return nil +} + +// RecordActivity updates the last-activity timestamp for a reserved device. +func (m *Manager) RecordActivity(platform, target string) { + key := sessionKey{platform: platform, target: target} + m.mu.Lock() + defer m.mu.Unlock() + if _, ok := m.reservations[key]; ok { + m.activity[key] = m.now() + } +} + +// Release removes a reservation and its activity tracking. +func (m *Manager) Release(platform, target string) { + key := sessionKey{platform: platform, target: target} + m.mu.Lock() + defer m.mu.Unlock() + delete(m.reservations, key) + delete(m.activity, key) +} + +// Owner returns the session that owns a device, if any. +func (m *Manager) Owner(platform, target string) (string, bool) { + key := sessionKey{platform: platform, target: target} + m.mu.Lock() + defer m.mu.Unlock() + owner, ok := m.reservations[key] + return owner, ok +} + +// Start runs the idle sweeper until ctx is cancelled. +func (m *Manager) Start(ctx context.Context) { + if m.timeout <= 0 { + return + } + + interval := m.timeout / 2 + if interval < time.Minute { + interval = time.Minute + } + + ticker := time.NewTicker(interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + m.sweep(ctx) + } + } +} + +type idleItem struct { + key sessionKey + owner string +} + +func (m *Manager) sweep(ctx context.Context) { + now := m.now() + m.mu.Lock() + var idle []idleItem + for key, t := range m.activity { + if now.Sub(t) > m.timeout { + idle = append(idle, idleItem{key: key, owner: m.reservations[key]}) + } + } + m.mu.Unlock() + + for _, it := range idle { + m.logger.Info("stopping idle device", + "platform", it.key.platform, + "target", it.key.target, + "owner", it.owner) + + if m.stopper != nil { + if err := m.stopper(ctx, it.key.platform, it.key.target, it.owner); err != nil { + m.logger.Error("failed to stop idle device", "platform", it.key.platform, "target", it.key.target, "error", err) + } + } + + m.mu.Lock() + if m.reservations[it.key] == it.owner { + delete(m.reservations, it.key) + delete(m.activity, it.key) + } + m.mu.Unlock() + } +} diff --git a/internal/core/session_test.go b/internal/core/session_test.go new file mode 100644 index 0000000..c39f790 --- /dev/null +++ b/internal/core/session_test.go @@ -0,0 +1,204 @@ +package core + +import ( + "context" + "errors" + "io" + "log/slog" + "testing" + "time" + + "github.com/espetro/mcp-sim/pkg/contract" +) + +func discardLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +func TestReserveAndOwner(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("Reserve failed: %v", err) + } + + owner, ok := m.Owner("ios", "sim1") + if !ok || owner != "session-a" { + t.Fatalf("Owner = (%q, %v), want (session-a, true)", owner, ok) + } + + if err := m.CheckAccess("ios", "sim1", "session-a"); err != nil { + t.Fatalf("CheckAccess by owner failed: %v", err) + } + + var te *ToolError + if err := m.CheckAccess("ios", "sim1", "session-b"); !errors.As(err, &te) || te.Code != contract.ErrDeviceReserved { + t.Fatalf("CheckAccess by other session returned %v, want device_reserved", err) + } +} + +func TestReserveIdempotent(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("first Reserve failed: %v", err) + } + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("second Reserve by same session failed: %v", err) + } + if err := m.Reserve(ctx, "ios", "sim1", "session-b"); err == nil { + t.Fatal("Reserve by different session should fail") + } +} + +func TestRelease(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("Reserve failed: %v", err) + } + m.Release("ios", "sim1") + + if _, ok := m.Owner("ios", "sim1"); ok { + t.Fatal("Owner should be false after Release") + } + if err := m.CheckAccess("ios", "sim1", "session-b"); err != nil { + t.Fatalf("CheckAccess after release failed: %v", err) + } +} + +func TestCheckAccessUnreserved(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + if err := m.CheckAccess("ios", "sim1", "any-session"); err != nil { + t.Fatalf("CheckAccess on unreserved device failed: %v", err) + } +} + +func TestRecordActivityOnlyWhenReserved(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + base := time.Date(2026, 7, 5, 0, 0, 0, 0, time.UTC) + m.now = func() time.Time { return base } + m.RecordActivity("ios", "sim1") + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("Reserve failed: %v", err) + } + m.now = func() time.Time { return base.Add(2 * time.Hour) } + m.RecordActivity("ios", "sim1") +} + +func TestSweeperStopsIdleAndReleases(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + stopped := make(map[sessionKey]bool) + m.SetStopper(func(_ context.Context, platform, target, owner string) error { + stopped[sessionKey{platform, target}] = true + if owner != "session-a" { + return errors.New("unexpected owner") + } + return nil + }) + + base := time.Date(2026, 7, 5, 0, 0, 0, 0, time.UTC) + m.now = func() time.Time { return base } + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("Reserve failed: %v", err) + } + + m.now = func() time.Time { return base.Add(2 * time.Hour) } + m.sweep(ctx) + + if !stopped[sessionKey{"ios", "sim1"}] { + t.Fatal("sweeper did not stop idle device") + } + if _, ok := m.Owner("ios", "sim1"); ok { + t.Fatal("reservation should be released after idle sweep") + } +} + +func TestSweeperDoesNotStopActive(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + stopped := false + m.SetStopper(func(_ context.Context, _, _, _ string) error { + stopped = true + return nil + }) + + base := time.Date(2026, 7, 5, 0, 0, 0, 0, time.UTC) + m.now = func() time.Time { return base } + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("Reserve failed: %v", err) + } + + m.now = func() time.Time { return base.Add(30 * time.Minute) } + m.sweep(ctx) + + if stopped { + t.Fatal("sweeper stopped active device") + } + if _, ok := m.Owner("ios", "sim1"); !ok { + t.Fatal("active reservation should remain") + } +} + +func TestSweeperDoesNotReleaseChangedOwner(t *testing.T) { + m := NewManager(time.Hour, discardLogger()) + ctx := context.Background() + + stopped := make(map[sessionKey]bool) + m.SetStopper(func(_ context.Context, platform, target, owner string) error { + stopped[sessionKey{platform, target}] = true + // Simulate another goroutine claiming the device while stopper runs. + key := sessionKey{platform, target} + m.mu.Lock() + m.reservations[key] = "session-c" + m.mu.Unlock() + return nil + }) + + base := time.Date(2026, 7, 5, 0, 0, 0, 0, time.UTC) + m.now = func() time.Time { return base } + + if err := m.Reserve(ctx, "ios", "sim1", "session-a"); err != nil { + t.Fatalf("Reserve failed: %v", err) + } + + m.now = func() time.Time { return base.Add(2 * time.Hour) } + m.sweep(ctx) + + if !stopped[sessionKey{"ios", "sim1"}] { + t.Fatal("sweeper should still attempt to stop the idle key") + } + owner, ok := m.Owner("ios", "sim1") + if !ok || owner != "session-c" { + t.Fatalf("Owner = (%q, %v), want (session-c, true)", owner, ok) + } +} + +func TestStartReturnsImmediatelyWhenDisabled(t *testing.T) { + m := NewManager(0, discardLogger()) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + done := make(chan struct{}) + go func() { + m.Start(ctx) + close(done) + }() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("Start did not return immediately with timeout=0") + } +} diff --git a/internal/core/tools.go b/internal/core/tools.go index a4961c4..cc2b96a 100644 --- a/internal/core/tools.go +++ b/internal/core/tools.go @@ -19,28 +19,37 @@ func (e *ToolError) Error() string { } // ListDevices returns all devices across all platforms. -func ListDevices(ctx context.Context, registry *Registry) ([]contract.Device, error) { +func ListDevices(ctx context.Context, registry *Registry, sessions *Manager) ([]contract.Device, error) { var result []contract.Device for name, p := range registry.AllPlatforms() { devs, err := p.List(ctx) if err != nil { return nil, fmt.Errorf("listing devices on %s: %w", name, err) } + for i := range devs { + if owner, ok := sessions.Owner(name, devs[i].ID); ok { + devs[i].OwnerSession = owner + } + } result = append(result, devs...) } return result, nil } // GetDeviceState returns the state of a specific device. -func GetDeviceState(ctx context.Context, registry *Registry, platformName, target string) (contract.DeviceState, error) { +func GetDeviceState(ctx context.Context, registry *Registry, sessions *Manager, platformName, target, sessionID string) (contract.DeviceState, error) { p, ok := registry.PlatformByName(platformName) if !ok { return contract.DeviceStateUnknown, &ToolError{Code: contract.ErrUnsupportedPlatform, Msg: "platform not found: " + platformName} } + if err := sessions.CheckAccess(platformName, target, sessionID); err != nil { + return contract.DeviceStateUnknown, err + } state, err := p.State(ctx, target) if err != nil { return contract.DeviceStateUnknown, err } + sessions.RecordActivity(platformName, target) return state, nil } @@ -72,10 +81,17 @@ func ControllerStatus(ctx context.Context, registry *Registry, name string) (con } // AwaitDeviceReady waits for a device to become ready. -func AwaitDeviceReady(ctx context.Context, registry *Registry, platformName, target string, timeout time.Duration) error { +func AwaitDeviceReady(ctx context.Context, registry *Registry, sessions *Manager, platformName, target, sessionID string, timeout time.Duration) error { p, ok := registry.PlatformByName(platformName) if !ok { return &ToolError{Code: contract.ErrUnsupportedPlatform, Msg: "platform not found: " + platformName} } - return p.AwaitReady(ctx, target, timeout) + if err := sessions.CheckAccess(platformName, target, sessionID); err != nil { + return err + } + if err := p.AwaitReady(ctx, target, timeout); err != nil { + return err + } + sessions.RecordActivity(platformName, target) + return nil } diff --git a/pkg/contract/types.go b/pkg/contract/types.go index 98702dc..bc6d596 100644 --- a/pkg/contract/types.go +++ b/pkg/contract/types.go @@ -23,6 +23,8 @@ type Device struct { State DeviceState `json:"state"` // OS version if known Version string `json:"version,omitempty"` + // OwnerSession is the session_id that currently reserves this device. + OwnerSession string `json:"owner_session,omitempty"` } // StartOpts controls how a device is started. @@ -52,5 +54,6 @@ const ( ErrTimeout = "timeout" ErrUnsupportedPlatform = "unsupported_platform" ErrUnsupportedController = "unsupported_controller" + ErrDeviceReserved = "device_reserved" ErrInternal = "internal" ) diff --git a/pkg/mcp/server.go b/pkg/mcp/server.go index 43bac45..4982ad1 100644 --- a/pkg/mcp/server.go +++ b/pkg/mcp/server.go @@ -19,7 +19,7 @@ type Server struct { } // NewServer creates an MCP server with the mcp-sim tool set registered. -func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog.Logger) *Server { +func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, sessions *core.Manager, logger *slog.Logger) *Server { s := mcp.NewServer(&mcp.Implementation{ Name: "mcp-sim", Title: "MCP Simulator Server", @@ -35,7 +35,7 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. }, func(ctx context.Context, _ *mcp.CallToolRequest, _ struct{}) (*mcp.CallToolResult, struct { Devices []contract.Device `json:"devices"` }, error) { - devs, err := core.ListDevices(ctx, registry) + devs, err := core.ListDevices(ctx, registry, sessions) return nil, struct { Devices []contract.Device `json:"devices"` }{Devices: devs}, err @@ -46,18 +46,19 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. Name: "boot_device", Description: "Boot a device by platform and target identifier.", }, func(ctx context.Context, _ *mcp.CallToolRequest, in struct { - Platform string `json:"platform"` - Target string `json:"target"` - NoWindow bool `json:"no_window,omitempty"` - Port int `json:"port,omitempty"` - Timeout int `json:"timeout,omitempty"` + Platform string `json:"platform"` + Target string `json:"target"` + NoWindow bool `json:"no_window,omitempty"` + Port int `json:"port,omitempty"` + Timeout int `json:"timeout,omitempty"` + SessionID string `json:"session_id,omitempty"` }) (*mcp.CallToolResult, contract.Device, error) { opts := contract.StartOpts{ NoWindow: in.NoWindow, Port: in.Port, Timeout: time.Duration(in.Timeout) * time.Second, } - dev, err := lifecycle.BootDevice(ctx, in.Platform, in.Target, opts) + dev, err := lifecycle.BootDevice(ctx, in.Platform, in.Target, in.SessionID, opts) return nil, dev, err }) @@ -66,13 +67,14 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. Name: "stop_device", Description: "Stop a running device by platform and target identifier.", }, func(ctx context.Context, _ *mcp.CallToolRequest, in struct { - Platform string `json:"platform"` - Target string `json:"target"` + Platform string `json:"platform"` + Target string `json:"target"` + SessionID string `json:"session_id,omitempty"` }) (*mcp.CallToolResult, contract.Device, error) { - if err := lifecycle.StopDevice(ctx, in.Platform, in.Target); err != nil { + if err := lifecycle.StopDevice(ctx, in.Platform, in.Target, in.SessionID); err != nil { return nil, contract.Device{}, err } - state, err := core.GetDeviceState(ctx, registry, in.Platform, in.Target) + state, err := core.GetDeviceState(ctx, registry, sessions, in.Platform, in.Target, in.SessionID) return nil, contract.Device{Platform: in.Platform, ID: in.Target, State: state}, err }) @@ -81,13 +83,14 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. Name: "wipe_device", Description: "Wipe a device, erasing its user data. Stops the device first if needed.", }, func(ctx context.Context, _ *mcp.CallToolRequest, in struct { - Platform string `json:"platform"` - Target string `json:"target"` + Platform string `json:"platform"` + Target string `json:"target"` + SessionID string `json:"session_id,omitempty"` }) (*mcp.CallToolResult, contract.Device, error) { - if err := lifecycle.WipeDevice(ctx, in.Platform, in.Target); err != nil { + if err := lifecycle.WipeDevice(ctx, in.Platform, in.Target, in.SessionID); err != nil { return nil, contract.Device{}, err } - state, err := core.GetDeviceState(ctx, registry, in.Platform, in.Target) + state, err := core.GetDeviceState(ctx, registry, sessions, in.Platform, in.Target, in.SessionID) return nil, contract.Device{Platform: in.Platform, ID: in.Target, State: state}, err }) @@ -96,12 +99,13 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. Name: "get_state", Description: "Get the current state of a device (stopped/booting/running/error).", }, func(ctx context.Context, _ *mcp.CallToolRequest, in struct { - Platform string `json:"platform"` - Target string `json:"target"` + Platform string `json:"platform"` + Target string `json:"target"` + SessionID string `json:"session_id,omitempty"` }) (*mcp.CallToolResult, struct { State string `json:"state"` }, error) { - state, err := core.GetDeviceState(ctx, registry, in.Platform, in.Target) + state, err := core.GetDeviceState(ctx, registry, sessions, in.Platform, in.Target, in.SessionID) return nil, struct { State string `json:"state"` }{State: string(state)}, err @@ -112,15 +116,16 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. Name: "await_ready", Description: "Block until the device is fully booted, or the timeout fires.", }, func(ctx context.Context, _ *mcp.CallToolRequest, in struct { - Platform string `json:"platform"` - Target string `json:"target"` - Timeout int `json:"timeout,omitempty"` + Platform string `json:"platform"` + Target string `json:"target"` + Timeout int `json:"timeout,omitempty"` + SessionID string `json:"session_id,omitempty"` }) (*mcp.CallToolResult, struct{ Ready bool }, error) { timeout := time.Duration(in.Timeout) * time.Second if timeout == 0 { timeout = 180 * time.Second } - if err := core.AwaitDeviceReady(ctx, registry, in.Platform, in.Target, timeout); err != nil { + if err := core.AwaitDeviceReady(ctx, registry, sessions, in.Platform, in.Target, in.SessionID, timeout); err != nil { return nil, struct{ Ready bool }{}, err } return nil, struct{ Ready bool }{Ready: true}, nil @@ -131,11 +136,12 @@ func NewServer(registry *core.Registry, lifecycle *core.Lifecycle, logger *slog. Name: "open_url", Description: "Open a URL or deep link on a device.", }, func(ctx context.Context, _ *mcp.CallToolRequest, in struct { - Platform string `json:"platform"` - Target string `json:"target"` - URL string `json:"url"` + Platform string `json:"platform"` + Target string `json:"target"` + URL string `json:"url"` + SessionID string `json:"session_id,omitempty"` }) (*mcp.CallToolResult, struct{ Success bool }, error) { - if err := lifecycle.OpenURL(ctx, in.Platform, in.Target, in.URL); err != nil { + if err := lifecycle.OpenURL(ctx, in.Platform, in.Target, in.SessionID, in.URL); err != nil { return nil, struct{ Success bool }{}, err } return nil, struct{ Success bool }{Success: true}, nil diff --git a/platforms/android/android.go b/platforms/android/android.go index 14d0a1d..1ecaa4f 100644 --- a/platforms/android/android.go +++ b/platforms/android/android.go @@ -148,7 +148,9 @@ func (p *Platform) List(ctx context.Context) ([]contract.Device, error) { var devs []contract.Device for _, name := range avds { + p.mu.Lock() port, ok := p.avdPortMap[name] + p.mu.Unlock() state := contract.DeviceStateStopped if ok { serial := fmt.Sprintf("emulator-%d", port) @@ -257,7 +259,9 @@ func (p *Platform) Stop(ctx context.Context, target string) error { // State returns the state of an emulator. func (p *Platform) State(ctx context.Context, target string) (contract.DeviceState, error) { + p.mu.Lock() port, ok := p.avdPortMap[target] + p.mu.Unlock() if !ok { // Try to discover port from adb devices. out, err := p.adbCmd(ctx, "devices").Output() @@ -277,7 +281,9 @@ func (p *Platform) State(ctx context.Context, target string) (contract.DeviceSta serial := parts[0] if strings.HasPrefix(serial, "emulator-") { if n, err := strconv.Atoi(strings.TrimPrefix(serial, "emulator-")); err == nil { + p.mu.Lock() p.avdPortMap[target] = n + p.mu.Unlock() port = n break }