Replace supervisor with agentrun worker service
Move agent-run orchestration from the legacy supervisor path into the new agentrun worker/service package and wire it through coredata, server, policies, and GraphQL resolvers. This consolidates run lifecycle handling around lease-aware workers and aligns API surface with the new agent-run domain model so reviewers can follow one coherent execution path. Signed-off-by: Bryan Frimin <bryan@probo.com>
This commit is contained in:
@@ -62,7 +62,7 @@ the cancellation so they complete naturally; the cancel is only
|
|||||||
observed at the next safe boundary.
|
observed at the next safe boundary.
|
||||||
|
|
||||||
Use `agent.ErrSuspendForCheckpoint` as the cancel cause when the
|
Use `agent.ErrSuspendForCheckpoint` as the cancel cause when the
|
||||||
intent is graceful suspend — supervisors that distinguish a
|
intent is graceful suspend — workers that distinguish a
|
||||||
graceful-stop request from infrastructure-level causes (lease loss,
|
graceful-stop request from infrastructure-level causes (lease loss,
|
||||||
heartbeat failure) inspect `context.Cause(ctx)` to dispatch.
|
heartbeat failure) inspect `context.Cause(ctx)` to dispatch.
|
||||||
|
|
||||||
@@ -77,7 +77,7 @@ Implications:
|
|||||||
it must derive its own with `context.WithTimeout(ctx, ...)` inside
|
it must derive its own with `context.WithTimeout(ctx, ...)` inside
|
||||||
the tool body.
|
the tool body.
|
||||||
|
|
||||||
The supervisor (`pkg/probo/agent_run_handler.go`) maps a SIGTERM-driven
|
The agent run worker (`pkg/agentrun/handler.go`) maps a SIGTERM-driven
|
||||||
shutdown broadcast onto a per-run `cancelRun(agent.ErrSuspendForCheckpoint)`,
|
shutdown broadcast onto a per-run `cancelRun(agent.ErrSuspendForCheckpoint)`,
|
||||||
so the same contract drives both the public Go API and the worker
|
so the same contract drives both the public Go API and the worker
|
||||||
infrastructure path.
|
infrastructure path.
|
||||||
|
|||||||
@@ -63,9 +63,9 @@ type (
|
|||||||
Result ToolResult
|
Result ToolResult
|
||||||
}
|
}
|
||||||
|
|
||||||
// Checkpointer is supervisor-internal. Implementations may use raw
|
// Checkpointer is worker-internal. Implementations may use raw
|
||||||
// run IDs because public API/service methods perform tenant scoping and
|
// run IDs because public API/service methods perform tenant scoping and
|
||||||
// authorization before a run reaches the supervisor.
|
// authorization before a run reaches the worker.
|
||||||
Checkpointer interface {
|
Checkpointer interface {
|
||||||
Save(ctx context.Context, runID string, cp *Checkpoint) error
|
Save(ctx context.Context, runID string, cp *Checkpoint) error
|
||||||
Load(ctx context.Context, runID string) (*Checkpoint, error)
|
Load(ctx context.Context, runID string) (*Checkpoint, error)
|
||||||
@@ -76,7 +76,7 @@ type (
|
|||||||
}
|
}
|
||||||
|
|
||||||
SuspendedError struct {
|
SuspendedError struct {
|
||||||
RunID string // Set when the outer loop has a store+runID (supervisor-managed).
|
RunID string // Set when the outer loop has a store+runID (worker-managed).
|
||||||
Checkpoint *Checkpoint // Set when returning from an inner agent-as-tool (no store).
|
Checkpoint *Checkpoint // Set when returning from an inner agent-as-tool (no store).
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
24
pkg/agentrun/export_test.go
Normal file
24
pkg/agentrun/export_test.go
Normal file
@@ -0,0 +1,24 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
// purpose with or without fee is hereby granted, provided that the above
|
||||||
|
// copyright notice and this permission notice appear in all copies.
|
||||||
|
//
|
||||||
|
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
package agentrun
|
||||||
|
|
||||||
|
// ShutdownBroadcast returns the channel that closes once the worker has
|
||||||
|
// broadcast graceful shutdown to all in-flight runs. It is compiled only
|
||||||
|
// in test builds so external tests can synchronize tool release with
|
||||||
|
// shutdown propagation without leaking a test-only method into the
|
||||||
|
// worker's public API.
|
||||||
|
func (w *Worker) ShutdownBroadcast() <-chan struct{} {
|
||||||
|
return w.handler.shutdownCh
|
||||||
|
}
|
||||||
@@ -12,7 +12,7 @@
|
|||||||
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
// PERFORMANCE OF THIS SOFTWARE.
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
package probo
|
package agentrun
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
@@ -31,30 +31,28 @@ import (
|
|||||||
"go.probo.inc/probo/pkg/llm"
|
"go.probo.inc/probo/pkg/llm"
|
||||||
)
|
)
|
||||||
|
|
||||||
type agentRunHandler struct {
|
type handler struct {
|
||||||
pg *pg.Client
|
pg *pg.Client
|
||||||
store *coredata.PGCheckpointer
|
store *coredata.PGCheckpointer
|
||||||
registry agent.AgentRegistry
|
registry agent.AgentRegistry
|
||||||
logger *log.Logger
|
logger *log.Logger
|
||||||
leaseDuration time.Duration
|
leaseDuration time.Duration
|
||||||
workerID string
|
|
||||||
shutdownCh chan struct{}
|
shutdownCh chan struct{}
|
||||||
shutdownOnce sync.Once
|
shutdownOnce sync.Once
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
_ worker.Handler[coredata.AgentRun] = (*agentRunHandler)(nil)
|
_ worker.Handler[coredata.AgentRun] = (*handler)(nil)
|
||||||
_ worker.StaleRecoverer = (*agentRunHandler)(nil)
|
_ worker.StaleRecoverer = (*handler)(nil)
|
||||||
)
|
)
|
||||||
|
|
||||||
// Claim loads the next pending agent run, marks it RUNNING with a lease
|
// Claim loads the next pending agent run, marks it RUNNING with a lease
|
||||||
// owned by this worker, and returns the row. When no work is available it
|
// owned by this worker, and returns the row. When no work is available it
|
||||||
// returns worker.ErrNoTask so the kit can back off until the next tick.
|
// returns worker.ErrNoTask so the kit can back off until the next tick.
|
||||||
func (h *agentRunHandler) Claim(ctx context.Context) (coredata.AgentRun, error) {
|
func (h *handler) Claim(ctx context.Context) (coredata.AgentRun, error) {
|
||||||
var (
|
var (
|
||||||
run = coredata.AgentRun{}
|
run = coredata.AgentRun{}
|
||||||
now = time.Now()
|
now = time.Now()
|
||||||
leaseOwner = h.workerID
|
|
||||||
leaseExpiresAt = now.Add(h.leaseDuration)
|
leaseExpiresAt = now.Add(h.leaseDuration)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -67,8 +65,8 @@ func (h *agentRunHandler) Claim(ctx context.Context) (coredata.AgentRun, error)
|
|||||||
|
|
||||||
run.Status = coredata.AgentRunStatusRunning
|
run.Status = coredata.AgentRunStatusRunning
|
||||||
run.StartedAt = &now
|
run.StartedAt = &now
|
||||||
run.LeaseOwner = &leaseOwner
|
|
||||||
run.LeaseExpiresAt = &leaseExpiresAt
|
run.LeaseExpiresAt = &leaseExpiresAt
|
||||||
|
run.LeaseGeneration++
|
||||||
run.UpdatedAt = now
|
run.UpdatedAt = now
|
||||||
|
|
||||||
if err := run.Update(ctx, tx, coredata.NewNoScope()); err != nil {
|
if err := run.Update(ctx, tx, coredata.NewNoScope()); err != nil {
|
||||||
@@ -99,9 +97,10 @@ func (h *agentRunHandler) Claim(ctx context.Context) (coredata.AgentRun, error)
|
|||||||
// nil is returned for both successful runs and graceful exits
|
// nil is returned for both successful runs and graceful exits
|
||||||
// (lease loss, infrastructure suspension) where the row state is
|
// (lease loss, infrastructure suspension) where the row state is
|
||||||
// already consistent.
|
// already consistent.
|
||||||
func (h *agentRunHandler) Process(ctx context.Context, run coredata.AgentRun) error {
|
func (h *handler) Process(ctx context.Context, run coredata.AgentRun) error {
|
||||||
runCtx, cancelRun := context.WithCancelCause(ctx)
|
runCtx, cancelRun := context.WithCancelCause(ctx)
|
||||||
defer cancelRun(nil)
|
defer cancelRun(nil)
|
||||||
|
leaseGeneration := run.LeaseGeneration
|
||||||
|
|
||||||
forwarderDone := make(chan struct{})
|
forwarderDone := make(chan struct{})
|
||||||
defer close(forwarderDone)
|
defer close(forwarderDone)
|
||||||
@@ -117,14 +116,14 @@ func (h *agentRunHandler) Process(ctx context.Context, run coredata.AgentRun) er
|
|||||||
heartbeatCtx, cancelHeartbeat := context.WithCancel(ctx)
|
heartbeatCtx, cancelHeartbeat := context.WithCancel(ctx)
|
||||||
defer cancelHeartbeat()
|
defer cancelHeartbeat()
|
||||||
|
|
||||||
go h.heartbeatLease(heartbeatCtx, run.ID.String(), cancelRun)
|
go h.heartbeatLease(heartbeatCtx, run.ID.String(), leaseGeneration, cancelRun)
|
||||||
|
|
||||||
return h.executeRun(runCtx, &run)
|
return h.executeRun(runCtx, &run, leaseGeneration)
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecoverStale resets agent runs whose worker lease has expired back to
|
// RecoverStale resets agent runs whose worker lease has expired back to
|
||||||
// PENDING so a fresh supervisor can pick them up on the next cycle.
|
// PENDING so a fresh worker can pick them up on the next cycle.
|
||||||
func (h *agentRunHandler) RecoverStale(ctx context.Context) error {
|
func (h *handler) RecoverStale(ctx context.Context) error {
|
||||||
if err := h.pg.WithConn(
|
if err := h.pg.WithConn(
|
||||||
ctx,
|
ctx,
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
@@ -141,13 +140,14 @@ func (h *agentRunHandler) RecoverStale(ctx context.Context) error {
|
|||||||
// in-flight Process forwarder goroutines observe the close and propagate
|
// in-flight Process forwarder goroutines observe the close and propagate
|
||||||
// it to their per-run agent stop channels, letting agents checkpoint at
|
// it to their per-run agent stop channels, letting agents checkpoint at
|
||||||
// the next turn boundary before Process returns.
|
// the next turn boundary before Process returns.
|
||||||
func (h *agentRunHandler) signalShutdown() {
|
func (h *handler) signalShutdown() {
|
||||||
h.shutdownOnce.Do(func() { close(h.shutdownCh) })
|
h.shutdownOnce.Do(func() { close(h.shutdownCh) })
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *agentRunHandler) heartbeatLease(
|
func (h *handler) heartbeatLease(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
runID string,
|
runID string,
|
||||||
|
leaseGeneration int64,
|
||||||
cancelRun context.CancelCauseFunc,
|
cancelRun context.CancelCauseFunc,
|
||||||
) {
|
) {
|
||||||
ticker := time.NewTicker(h.leaseDuration / 3)
|
ticker := time.NewTicker(h.leaseDuration / 3)
|
||||||
@@ -163,13 +163,19 @@ func (h *agentRunHandler) heartbeatLease(
|
|||||||
if err := h.pg.WithConn(
|
if err := h.pg.WithConn(
|
||||||
ctx,
|
ctx,
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
rowsAffected, err := coredata.HeartbeatAgentRunLease(ctx, conn, runID, h.workerID, expiresAt)
|
rowsAffected, err := coredata.HeartbeatAgentRunLease(
|
||||||
|
ctx,
|
||||||
|
conn,
|
||||||
|
runID,
|
||||||
|
leaseGeneration,
|
||||||
|
expiresAt,
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if rowsAffected == 0 {
|
if rowsAffected == 0 {
|
||||||
return ErrAgentRunLeaseLost
|
return ErrLeaseLost
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -177,10 +183,10 @@ func (h *agentRunHandler) heartbeatLease(
|
|||||||
); err != nil {
|
); err != nil {
|
||||||
h.logger.ErrorCtx(ctx, "cannot heartbeat agent run lease", log.Error(err))
|
h.logger.ErrorCtx(ctx, "cannot heartbeat agent run lease", log.Error(err))
|
||||||
|
|
||||||
if errors.Is(err, ErrAgentRunLeaseLost) {
|
if errors.Is(err, ErrLeaseLost) {
|
||||||
cancelRun(ErrAgentRunLeaseLost)
|
cancelRun(ErrLeaseLost)
|
||||||
} else {
|
} else {
|
||||||
cancelRun(fmt.Errorf("%w: %w", ErrAgentRunHeartbeatFailed, err))
|
cancelRun(fmt.Errorf("%w: %w", ErrHeartbeatFailed, err))
|
||||||
}
|
}
|
||||||
|
|
||||||
return
|
return
|
||||||
@@ -190,21 +196,21 @@ func (h *agentRunHandler) heartbeatLease(
|
|||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// agentRunErrorMessageMaxLen caps the error string persisted to the
|
// errorMessageMaxLen caps the error string persisted to the
|
||||||
// agent_runs.error_message column. Raw tool or LLM errors can embed
|
// agent_runs.error_message column. Raw tool or LLM errors can embed
|
||||||
// URLs with credentials, response snippets containing PII, or partial
|
// URLs with credentials, response snippets containing PII, or partial
|
||||||
// records from failed DB lookups; the full context is logged while
|
// records from failed DB lookups; the full context is logged while
|
||||||
// only a truncated summary is stored for caller-visible state.
|
// only a truncated summary is stored for caller-visible state.
|
||||||
agentRunErrorMessageMaxLen = 512
|
errorMessageMaxLen = 512
|
||||||
)
|
)
|
||||||
|
|
||||||
func sanitizeAgentRunError(err error) string {
|
func sanitizeError(err error) string {
|
||||||
msg := err.Error()
|
msg := err.Error()
|
||||||
if len(msg) <= agentRunErrorMessageMaxLen {
|
if len(msg) <= errorMessageMaxLen {
|
||||||
return msg
|
return msg
|
||||||
}
|
}
|
||||||
|
|
||||||
cut := agentRunErrorMessageMaxLen
|
cut := errorMessageMaxLen
|
||||||
for cut > 0 && !utf8.RuneStart(msg[cut]) {
|
for cut > 0 && !utf8.RuneStart(msg[cut]) {
|
||||||
cut--
|
cut--
|
||||||
}
|
}
|
||||||
@@ -212,8 +218,29 @@ func sanitizeAgentRunError(err error) string {
|
|||||||
return msg[:cut] + "…"
|
return msg[:cut] + "…"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun) error {
|
type leasedCheckpointer struct {
|
||||||
|
store *coredata.PGCheckpointer
|
||||||
|
leaseGeneration int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s leasedCheckpointer) Save(ctx context.Context, runID string, cp *agent.Checkpoint) error {
|
||||||
|
return s.store.SaveForLease(ctx, runID, cp, s.leaseGeneration)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s leasedCheckpointer) Load(ctx context.Context, runID string) (*agent.Checkpoint, error) {
|
||||||
|
return s.store.Load(ctx, runID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *handler) executeRun(
|
||||||
|
ctx context.Context,
|
||||||
|
run *coredata.AgentRun,
|
||||||
|
leaseGeneration int64,
|
||||||
|
) error {
|
||||||
runID := run.ID.String()
|
runID := run.ID.String()
|
||||||
|
checkpointer := leasedCheckpointer{
|
||||||
|
store: h.store,
|
||||||
|
leaseGeneration: leaseGeneration,
|
||||||
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
result *agent.Result
|
result *agent.Result
|
||||||
@@ -222,7 +249,7 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
|
|
||||||
if run.Checkpoint != nil {
|
if run.Checkpoint != nil {
|
||||||
h.logger.InfoCtx(ctx, "resuming agent run", log.String("run_id", runID))
|
h.logger.InfoCtx(ctx, "resuming agent run", log.String("run_id", runID))
|
||||||
result, runErr = agent.Restore(ctx, h.store, runID, h.registry)
|
result, runErr = agent.Restore(ctx, checkpointer, runID, h.registry)
|
||||||
} else {
|
} else {
|
||||||
h.logger.InfoCtx(ctx, "starting agent run", log.String("run_id", runID))
|
h.logger.InfoCtx(ctx, "starting agent run", log.String("run_id", runID))
|
||||||
|
|
||||||
@@ -237,7 +264,7 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
result, runErr = a.Run(
|
result, runErr = a.Run(
|
||||||
ctx,
|
ctx,
|
||||||
inputMsgs,
|
inputMsgs,
|
||||||
agent.WithCheckpointer(h.store, runID),
|
agent.WithCheckpointer(checkpointer, runID),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -246,7 +273,7 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
// Heartbeat loss: another worker may have taken over. Do not commit
|
// Heartbeat loss: another worker may have taken over. Do not commit
|
||||||
// any status — stale recovery will handle the row. Surface the cause
|
// any status — stale recovery will handle the row. Surface the cause
|
||||||
// so the worker kit logs and traces a failure for this attempt.
|
// so the worker kit logs and traces a failure for this attempt.
|
||||||
if cause := context.Cause(ctx); errors.Is(cause, ErrAgentRunLeaseLost) || errors.Is(cause, ErrAgentRunHeartbeatFailed) {
|
if cause := context.Cause(ctx); errors.Is(cause, ErrLeaseLost) || errors.Is(cause, ErrHeartbeatFailed) {
|
||||||
h.logger.WarnCtx(
|
h.logger.WarnCtx(
|
||||||
context.WithoutCancel(ctx),
|
context.WithoutCancel(ctx),
|
||||||
"agent run stopped after heartbeat failure; leaving status for stale recovery",
|
"agent run stopped after heartbeat failure; leaving status for stale recovery",
|
||||||
@@ -277,7 +304,6 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
now := time.Now()
|
now := time.Now()
|
||||||
run.UpdatedAt = now
|
run.UpdatedAt = now
|
||||||
run.StartedAt = nil
|
run.StartedAt = nil
|
||||||
run.LeaseOwner = nil
|
|
||||||
run.LeaseExpiresAt = nil
|
run.LeaseExpiresAt = nil
|
||||||
|
|
||||||
if runErr == nil {
|
if runErr == nil {
|
||||||
@@ -304,7 +330,7 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
log.String("run_id", runID),
|
log.String("run_id", runID),
|
||||||
log.Error(runErr),
|
log.Error(runErr),
|
||||||
)
|
)
|
||||||
msg := sanitizeAgentRunError(runErr)
|
msg := sanitizeError(runErr)
|
||||||
run.ErrorMessage = &msg
|
run.ErrorMessage = &msg
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -313,10 +339,15 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
if err := h.pg.WithTx(
|
if err := h.pg.WithTx(
|
||||||
commitCtx,
|
commitCtx,
|
||||||
func(ctx context.Context, tx pg.Tx) error {
|
func(ctx context.Context, tx pg.Tx) error {
|
||||||
if err := run.Update(ctx, tx, coredata.NewNoScope()); err != nil {
|
rowsAffected, err := coredata.CommitAgentRunResult(ctx, tx, run, leaseGeneration)
|
||||||
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if rowsAffected == 0 {
|
||||||
|
return ErrLeaseLost
|
||||||
|
}
|
||||||
|
|
||||||
if run.Status == coredata.AgentRunStatusCompleted {
|
if run.Status == coredata.AgentRunStatusCompleted {
|
||||||
if err := run.ClearCheckpoint(ctx, tx, coredata.NewNoScope()); err != nil {
|
if err := run.ClearCheckpoint(ctx, tx, coredata.NewNoScope()); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -326,6 +357,16 @@ func (h *agentRunHandler) executeRun(ctx context.Context, run *coredata.AgentRun
|
|||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
); err != nil {
|
); err != nil {
|
||||||
|
if errors.Is(err, ErrLeaseLost) {
|
||||||
|
h.logger.WarnCtx(
|
||||||
|
commitCtx,
|
||||||
|
"agent run lost lease before commit; discarding stale completion",
|
||||||
|
log.String("run_id", runID),
|
||||||
|
)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
h.logger.ErrorCtx(commitCtx, "cannot commit agent run status", log.Error(err))
|
h.logger.ErrorCtx(commitCtx, "cannot commit agent run status", log.Error(err))
|
||||||
return fmt.Errorf("cannot commit agent run status: %w", err)
|
return fmt.Errorf("cannot commit agent run status: %w", err)
|
||||||
}
|
}
|
||||||
111
pkg/agentrun/internal_test.go
Normal file
111
pkg/agentrun/internal_test.go
Normal file
@@ -0,0 +1,111 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
// purpose with or without fee is hereby granted, provided that the above
|
||||||
|
// copyright notice and this permission notice appear in all copies.
|
||||||
|
//
|
||||||
|
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
package agentrun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSanitizeError(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"short message unchanged",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
err := errors.New("short")
|
||||||
|
assert.Equal(t, "short", sanitizeError(err))
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"boundary length unchanged",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
msg := strings.Repeat("a", errorMessageMaxLen)
|
||||||
|
assert.Equal(t, msg, sanitizeError(errors.New(msg)))
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"long utf8 message is rune safe and suffixed",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
msg := strings.Repeat("é", errorMessageMaxLen)
|
||||||
|
sanitized := sanitizeError(errors.New(msg))
|
||||||
|
|
||||||
|
assert.True(t, strings.HasSuffix(sanitized, "…"))
|
||||||
|
assert.True(t, len(sanitized) <= errorMessageMaxLen+len("…"))
|
||||||
|
assert.True(t, strings.HasPrefix(msg, strings.TrimSuffix(sanitized, "…")))
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWorkerOptions(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"interval updates only when positive",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := workerConfig{interval: 3 * time.Second}
|
||||||
|
|
||||||
|
WithWorkerInterval(0)(&cfg)
|
||||||
|
assert.Equal(t, 3*time.Second, cfg.interval)
|
||||||
|
|
||||||
|
WithWorkerInterval(7 * time.Second)(&cfg)
|
||||||
|
assert.Equal(t, 7*time.Second, cfg.interval)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"lease duration updates only when positive",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := workerConfig{leaseDuration: 5 * time.Second}
|
||||||
|
|
||||||
|
WithWorkerLeaseDuration(-1)(&cfg)
|
||||||
|
assert.Equal(t, 5*time.Second, cfg.leaseDuration)
|
||||||
|
|
||||||
|
WithWorkerLeaseDuration(12 * time.Second)(&cfg)
|
||||||
|
assert.Equal(t, 12*time.Second, cfg.leaseDuration)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
t.Run(
|
||||||
|
"max concurrency updates only when positive",
|
||||||
|
func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := workerConfig{maxConcurrency: 2}
|
||||||
|
|
||||||
|
WithWorkerMaxConcurrency(0)(&cfg)
|
||||||
|
assert.Equal(t, 2, cfg.maxConcurrency)
|
||||||
|
|
||||||
|
WithWorkerMaxConcurrency(9)(&cfg)
|
||||||
|
assert.Equal(t, 9, cfg.maxConcurrency)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
114
pkg/agentrun/service.go
Normal file
114
pkg/agentrun/service.go
Normal file
@@ -0,0 +1,114 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
// purpose with or without fee is hereby granted, provided that the above
|
||||||
|
// copyright notice and this permission notice appear in all copies.
|
||||||
|
//
|
||||||
|
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
package agentrun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"go.gearno.de/kit/pg"
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/gid"
|
||||||
|
"go.probo.inc/probo/pkg/page"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Service struct {
|
||||||
|
pg *pg.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewService(pgClient *pg.Client) *Service {
|
||||||
|
return &Service{pg: pgClient}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) Get(
|
||||||
|
ctx context.Context,
|
||||||
|
scope coredata.Scoper,
|
||||||
|
agentRunID gid.GID,
|
||||||
|
) (*coredata.AgentRun, error) {
|
||||||
|
run := &coredata.AgentRun{}
|
||||||
|
|
||||||
|
err := s.pg.WithConn(
|
||||||
|
ctx,
|
||||||
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
if err := run.LoadByID(ctx, conn, scope, agentRunID); err != nil {
|
||||||
|
return fmt.Errorf("cannot load agent run: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return run, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) ListForOrganizationID(
|
||||||
|
ctx context.Context,
|
||||||
|
scope coredata.Scoper,
|
||||||
|
organizationID gid.GID,
|
||||||
|
cursor *page.Cursor[coredata.AgentRunOrderField],
|
||||||
|
) (*page.Page[*coredata.AgentRun, coredata.AgentRunOrderField], error) {
|
||||||
|
var runs coredata.AgentRuns
|
||||||
|
|
||||||
|
err := s.pg.WithConn(
|
||||||
|
ctx,
|
||||||
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
organization := &coredata.Organization{}
|
||||||
|
if err := organization.LoadByID(ctx, conn, scope, organizationID); err != nil {
|
||||||
|
return fmt.Errorf("cannot load organization: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := runs.LoadByOrganizationID(ctx, conn, scope, organization.ID, cursor); err != nil {
|
||||||
|
return fmt.Errorf("cannot load agent runs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return page.NewPage(runs, cursor), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) CountForOrganizationID(
|
||||||
|
ctx context.Context,
|
||||||
|
scope coredata.Scoper,
|
||||||
|
organizationID gid.GID,
|
||||||
|
) (int, error) {
|
||||||
|
var count int
|
||||||
|
|
||||||
|
err := s.pg.WithConn(
|
||||||
|
ctx,
|
||||||
|
func(ctx context.Context, conn pg.Querier) (err error) {
|
||||||
|
runs := &coredata.AgentRuns{}
|
||||||
|
|
||||||
|
count, err = runs.CountByOrganizationID(ctx, conn, scope, organizationID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("cannot count agent runs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
95
pkg/agentrun/service_test.go
Normal file
95
pkg/agentrun/service_test.go
Normal file
@@ -0,0 +1,95 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
// purpose with or without fee is hereby granted, provided that the above
|
||||||
|
// copyright notice and this permission notice appear in all copies.
|
||||||
|
//
|
||||||
|
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
package agentrun_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.probo.inc/probo/pkg/agentrun"
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/gid"
|
||||||
|
"go.probo.inc/probo/pkg/page"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestService_Get(t *testing.T) {
|
||||||
|
client := pgClient(t)
|
||||||
|
svc := agentrun.NewService(client)
|
||||||
|
|
||||||
|
run := insertPendingRun(
|
||||||
|
t,
|
||||||
|
client,
|
||||||
|
"service-get-agent",
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
got, err := svc.Get(context.Background(), coredata.NewNoScope(), run.ID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, got)
|
||||||
|
assert.Equal(t, run.ID, got.ID)
|
||||||
|
|
||||||
|
missingID := gid.New(run.ID.TenantID(), coredata.AgentRunEntityType)
|
||||||
|
_, err = svc.Get(context.Background(), coredata.NewNoScope(), missingID)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.ErrorIs(t, err, coredata.ErrResourceNotFound)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestService_ListForOrganizationID(t *testing.T) {
|
||||||
|
client := pgClient(t)
|
||||||
|
svc := agentrun.NewService(client)
|
||||||
|
|
||||||
|
orgID := insertTestOrganization(t, client)
|
||||||
|
|
||||||
|
runA := insertPendingRunInOrg(t, client, orgID, "service-list-agent-a", nil)
|
||||||
|
runB := insertPendingRunInOrg(t, client, orgID, "service-list-agent-b", nil)
|
||||||
|
|
||||||
|
cursor := page.NewCursor(
|
||||||
|
10,
|
||||||
|
nil,
|
||||||
|
page.Head,
|
||||||
|
page.OrderBy[coredata.AgentRunOrderField]{
|
||||||
|
Field: coredata.AgentRunOrderFieldCreatedAt,
|
||||||
|
Direction: page.OrderDirectionDesc,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
got, err := svc.ListForOrganizationID(context.Background(), coredata.NewNoScope(), orgID, cursor)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, got)
|
||||||
|
|
||||||
|
ids := make(map[gid.GID]bool)
|
||||||
|
for _, run := range got.Data {
|
||||||
|
ids[run.ID] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.True(t, ids[runA.ID])
|
||||||
|
assert.True(t, ids[runB.ID])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestService_CountForOrganizationID(t *testing.T) {
|
||||||
|
client := pgClient(t)
|
||||||
|
svc := agentrun.NewService(client)
|
||||||
|
|
||||||
|
orgID := insertTestOrganization(t, client)
|
||||||
|
_ = insertPendingRunInOrg(t, client, orgID, "service-count-agent-a", nil)
|
||||||
|
_ = insertPendingRunInOrg(t, client, orgID, "service-count-agent-b", nil)
|
||||||
|
_ = insertPendingRunInOrg(t, client, orgID, "service-count-agent-c", nil)
|
||||||
|
|
||||||
|
count, err := svc.CountForOrganizationID(context.Background(), coredata.NewNoScope(), orgID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, 3, count)
|
||||||
|
}
|
||||||
@@ -12,14 +12,13 @@
|
|||||||
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
// PERFORMANCE OF THIS SOFTWARE.
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
package probo
|
package agentrun
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go.gearno.de/crypto/uuid"
|
|
||||||
"go.gearno.de/kit/log"
|
"go.gearno.de/kit/log"
|
||||||
"go.gearno.de/kit/pg"
|
"go.gearno.de/kit/pg"
|
||||||
"go.gearno.de/kit/worker"
|
"go.gearno.de/kit/worker"
|
||||||
@@ -28,14 +27,14 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type (
|
type (
|
||||||
AgentRunSupervisor struct {
|
Worker struct {
|
||||||
handler *agentRunHandler
|
handler *handler
|
||||||
worker *worker.Worker[coredata.AgentRun]
|
kitWorker *worker.Worker[coredata.AgentRun]
|
||||||
}
|
}
|
||||||
|
|
||||||
AgentRunSupervisorOption func(*agentRunSupervisorConfig)
|
WorkerOption func(*workerConfig)
|
||||||
|
|
||||||
agentRunSupervisorConfig struct {
|
workerConfig struct {
|
||||||
interval time.Duration
|
interval time.Duration
|
||||||
leaseDuration time.Duration
|
leaseDuration time.Duration
|
||||||
maxConcurrency int
|
maxConcurrency int
|
||||||
@@ -43,42 +42,42 @@ type (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
ErrAgentRunHeartbeatFailed = errors.New("agent run heartbeat failed")
|
ErrHeartbeatFailed = errors.New("agent run heartbeat failed")
|
||||||
ErrAgentRunLeaseLost = errors.New("agent run lease lost")
|
ErrLeaseLost = errors.New("agent run lease lost")
|
||||||
)
|
)
|
||||||
|
|
||||||
func WithAgentRunSupervisorInterval(d time.Duration) AgentRunSupervisorOption {
|
func WithWorkerInterval(d time.Duration) WorkerOption {
|
||||||
return func(c *agentRunSupervisorConfig) {
|
return func(c *workerConfig) {
|
||||||
if d > 0 {
|
if d > 0 {
|
||||||
c.interval = d
|
c.interval = d
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func WithAgentRunSupervisorLeaseDuration(d time.Duration) AgentRunSupervisorOption {
|
func WithWorkerLeaseDuration(d time.Duration) WorkerOption {
|
||||||
return func(c *agentRunSupervisorConfig) {
|
return func(c *workerConfig) {
|
||||||
if d > 0 {
|
if d > 0 {
|
||||||
c.leaseDuration = d
|
c.leaseDuration = d
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func WithAgentRunSupervisorMaxConcurrency(n int) AgentRunSupervisorOption {
|
func WithWorkerMaxConcurrency(n int) WorkerOption {
|
||||||
return func(c *agentRunSupervisorConfig) {
|
return func(c *workerConfig) {
|
||||||
if n > 0 {
|
if n > 0 {
|
||||||
c.maxConcurrency = n
|
c.maxConcurrency = n
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewAgentRunSupervisor(
|
func NewWorker(
|
||||||
pgClient *pg.Client,
|
pgClient *pg.Client,
|
||||||
store *coredata.PGCheckpointer,
|
store *coredata.PGCheckpointer,
|
||||||
registry agent.AgentRegistry,
|
registry agent.AgentRegistry,
|
||||||
logger *log.Logger,
|
logger *log.Logger,
|
||||||
opts ...AgentRunSupervisorOption,
|
opts ...WorkerOption,
|
||||||
) *AgentRunSupervisor {
|
) *Worker {
|
||||||
cfg := agentRunSupervisorConfig{
|
cfg := workerConfig{
|
||||||
interval: 10 * time.Second,
|
interval: 10 * time.Second,
|
||||||
leaseDuration: 5 * time.Minute,
|
leaseDuration: 5 * time.Minute,
|
||||||
maxConcurrency: 5,
|
maxConcurrency: 5,
|
||||||
@@ -88,45 +87,35 @@ func NewAgentRunSupervisor(
|
|||||||
opt(&cfg)
|
opt(&cfg)
|
||||||
}
|
}
|
||||||
|
|
||||||
h := &agentRunHandler{
|
h := &handler{
|
||||||
pg: pgClient,
|
pg: pgClient,
|
||||||
store: store,
|
store: store,
|
||||||
registry: registry,
|
registry: registry,
|
||||||
logger: logger,
|
logger: logger,
|
||||||
leaseDuration: cfg.leaseDuration,
|
leaseDuration: cfg.leaseDuration,
|
||||||
workerID: uuid.MustNewV4().String(),
|
|
||||||
shutdownCh: make(chan struct{}),
|
shutdownCh: make(chan struct{}),
|
||||||
}
|
}
|
||||||
|
|
||||||
w := worker.New(
|
w := worker.New(
|
||||||
"agent-run-supervisor",
|
"agent-run-worker",
|
||||||
h,
|
h,
|
||||||
logger,
|
logger,
|
||||||
worker.WithInterval(cfg.interval),
|
worker.WithInterval(cfg.interval),
|
||||||
worker.WithMaxConcurrency(cfg.maxConcurrency),
|
worker.WithMaxConcurrency(cfg.maxConcurrency),
|
||||||
)
|
)
|
||||||
|
|
||||||
return &AgentRunSupervisor{handler: h, worker: w}
|
return &Worker{handler: h, kitWorker: w}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run starts the supervisor loop. It blocks until ctx is cancelled, then
|
// Run starts the worker loop. It blocks until ctx is cancelled, then
|
||||||
// closes the shutdown broadcast channel so in-flight Process calls can
|
// closes the shutdown broadcast channel so in-flight Process calls can
|
||||||
// checkpoint and exit, and waits for all of them to drain before
|
// checkpoint and exit, and waits for all of them to drain before
|
||||||
// returning.
|
// returning.
|
||||||
//
|
//
|
||||||
// signalShutdown is registered without a stop hook because it is
|
// signalShutdown is registered without a stop hook because it is
|
||||||
// idempotent (sync.Once) and we want it to fire on every ctx
|
// idempotent (sync.Once) and we want it to fire on every ctx
|
||||||
// cancellation, even one that races with worker.Run returning.
|
// cancellation, even one that races with kitWorker.Run returning.
|
||||||
func (s *AgentRunSupervisor) Run(ctx context.Context) error {
|
func (w *Worker) Run(ctx context.Context) error {
|
||||||
context.AfterFunc(ctx, s.handler.signalShutdown)
|
context.AfterFunc(ctx, w.handler.signalShutdown)
|
||||||
return s.worker.Run(ctx)
|
return w.kitWorker.Run(ctx)
|
||||||
}
|
|
||||||
|
|
||||||
// ShutdownBroadcastForTests returns a channel that closes once the
|
|
||||||
// supervisor has broadcast graceful shutdown to all in-flight runs.
|
|
||||||
// Exposed for tests that need to synchronize tool release with shutdown
|
|
||||||
// propagation; not a stable API and not part of the supervisor's public
|
|
||||||
// operational contract.
|
|
||||||
func (s *AgentRunSupervisor) ShutdownBroadcastForTests() <-chan struct{} {
|
|
||||||
return s.handler.shutdownCh
|
|
||||||
}
|
}
|
||||||
@@ -1,906 +0,0 @@
|
|||||||
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
|
||||||
//
|
|
||||||
// Permission to use, copy, modify, and/or distribute this software for any
|
|
||||||
// purpose with or without fee is hereby granted, provided that the above
|
|
||||||
// copyright notice and this permission notice appear in all copies.
|
|
||||||
//
|
|
||||||
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
|
||||||
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
|
||||||
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
|
||||||
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
|
||||||
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
|
||||||
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
|
||||||
// PERFORMANCE OF THIS SOFTWARE.
|
|
||||||
|
|
||||||
package agentruntest_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"os/exec"
|
|
||||||
"os/signal"
|
|
||||||
"path/filepath"
|
|
||||||
"strconv"
|
|
||||||
"sync"
|
|
||||||
"syscall"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"go.gearno.de/kit/log"
|
|
||||||
"go.gearno.de/kit/pg"
|
|
||||||
"go.probo.inc/probo/pkg/agent"
|
|
||||||
"go.probo.inc/probo/pkg/agentruntest"
|
|
||||||
"go.probo.inc/probo/pkg/coredata"
|
|
||||||
"go.probo.inc/probo/pkg/llm"
|
|
||||||
"go.probo.inc/probo/pkg/probo"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Test helpers
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func testLogger() *log.Logger {
|
|
||||||
return log.NewLogger(log.WithFormat(log.FormatPretty))
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Mock LLM provider
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
type mockProvider struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
responses []*llm.ChatCompletionResponse
|
|
||||||
calls int
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *mockProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
|
|
||||||
m.mu.Lock()
|
|
||||||
defer m.mu.Unlock()
|
|
||||||
|
|
||||||
if m.calls >= len(m.responses) {
|
|
||||||
return nil, errors.New("no more mock responses")
|
|
||||||
}
|
|
||||||
|
|
||||||
resp := m.responses[m.calls]
|
|
||||||
m.calls++
|
|
||||||
|
|
||||||
return resp, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *mockProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
|
|
||||||
return nil, errors.New("not implemented")
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestClient(provider llm.Provider) *llm.Client {
|
|
||||||
return llm.NewClient(provider, "test")
|
|
||||||
}
|
|
||||||
|
|
||||||
func stopResponse(text string) *llm.ChatCompletionResponse {
|
|
||||||
return &llm.ChatCompletionResponse{
|
|
||||||
Model: "test-model",
|
|
||||||
Message: llm.Message{
|
|
||||||
Role: llm.RoleAssistant,
|
|
||||||
Parts: []llm.Part{llm.TextPart{Text: text}},
|
|
||||||
},
|
|
||||||
Usage: llm.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
FinishReason: llm.FinishReasonStop,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func toolCallResponse(toolCalls ...llm.ToolCall) *llm.ChatCompletionResponse {
|
|
||||||
return &llm.ChatCompletionResponse{
|
|
||||||
Model: "test-model",
|
|
||||||
Message: llm.Message{
|
|
||||||
Role: llm.RoleAssistant,
|
|
||||||
ToolCalls: toolCalls,
|
|
||||||
},
|
|
||||||
Usage: llm.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
FinishReason: llm.FinishReasonToolCalls,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Simple agent registry
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
type simpleRegistry struct {
|
|
||||||
agents map[string]*agent.Agent
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *simpleRegistry) Agent(name string) (*agent.Agent, error) {
|
|
||||||
a, ok := r.agents[name]
|
|
||||||
if !ok {
|
|
||||||
return nil, fmt.Errorf("agent %q not found", name)
|
|
||||||
}
|
|
||||||
|
|
||||||
return a, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Test 3: Supervisor picks up a PENDING run and completes it
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
// Supervisor tests are intentionally sequential. The supervisor claims
|
|
||||||
// runs cross-tenant via LoadNextPendingForUpdateSkipLocked; running two
|
|
||||||
// supervisors against the same test database would steal each other's
|
|
||||||
// runs. If a per-tenant claim filter is ever added, these can go back
|
|
||||||
// to t.Parallel().
|
|
||||||
func TestAgentRunSupervisor_PicksUpAndCompletes(t *testing.T) {
|
|
||||||
client := agentruntest.PGClient(t)
|
|
||||||
store := coredata.NewPGCheckpointer(client)
|
|
||||||
|
|
||||||
provider := &mockProvider{
|
|
||||||
responses: []*llm.ChatCompletionResponse{
|
|
||||||
stopResponse("Done."),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
ag := agent.New(
|
|
||||||
"echo-agent",
|
|
||||||
newTestClient(provider),
|
|
||||||
agent.WithModel("test-model"),
|
|
||||||
agent.WithInstructions("Reply with done."),
|
|
||||||
)
|
|
||||||
|
|
||||||
registry := &simpleRegistry{
|
|
||||||
agents: map[string]*agent.Agent{"echo-agent": ag},
|
|
||||||
}
|
|
||||||
|
|
||||||
run := agentruntest.InsertPendingRun(
|
|
||||||
t,
|
|
||||||
client,
|
|
||||||
"echo-agent",
|
|
||||||
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "go"}}}},
|
|
||||||
)
|
|
||||||
|
|
||||||
supervisor := probo.NewAgentRunSupervisor(
|
|
||||||
client,
|
|
||||||
store,
|
|
||||||
registry,
|
|
||||||
testLogger(),
|
|
||||||
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
|
||||||
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
|
||||||
)
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
go func() { _ = supervisor.Run(ctx) }()
|
|
||||||
|
|
||||||
// Poll until the run is completed.
|
|
||||||
require.Eventually(
|
|
||||||
t,
|
|
||||||
func() bool {
|
|
||||||
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
|
||||||
return err == nil && r.Status == coredata.AgentRunStatusCompleted
|
|
||||||
},
|
|
||||||
10*time.Second,
|
|
||||||
200*time.Millisecond,
|
|
||||||
"run should reach COMPLETED status",
|
|
||||||
)
|
|
||||||
|
|
||||||
completed := agentruntest.LoadAgentRun(t, client, run.ID)
|
|
||||||
assert.Equal(t, coredata.AgentRunStatusCompleted, completed.Status)
|
|
||||||
assert.NotNil(t, completed.Result)
|
|
||||||
assert.Nil(t, completed.Checkpoint, "checkpoint should be cleared after completion")
|
|
||||||
assert.Nil(t, completed.ErrorMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Test 4: Supervisor stop/resume cycle with checkpoint
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func TestAgentRunSupervisor_StopAndResume(t *testing.T) {
|
|
||||||
client := agentruntest.PGClient(t)
|
|
||||||
store := coredata.NewPGCheckpointer(client)
|
|
||||||
|
|
||||||
// The tool blocks until signaled, giving us time to trigger graceful
|
|
||||||
// shutdown via the supervisor context while the agent is mid-turn.
|
|
||||||
toolReady := make(chan struct{})
|
|
||||||
toolRelease := make(chan struct{})
|
|
||||||
|
|
||||||
slowTool := agent.FunctionTool[struct{}](
|
|
||||||
"slow_work",
|
|
||||||
"Does slow work",
|
|
||||||
func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
|
|
||||||
close(toolReady)
|
|
||||||
<-toolRelease
|
|
||||||
|
|
||||||
return agent.ToolResult{Content: "work done"}, nil
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
// Provider sequence:
|
|
||||||
// Call 1: request tool call (first execution)
|
|
||||||
// Call 2: final stop response (after restoration)
|
|
||||||
provider := &mockProvider{
|
|
||||||
responses: []*llm.ChatCompletionResponse{
|
|
||||||
// First execution: LLM asks to call the tool.
|
|
||||||
toolCallResponse(llm.ToolCall{
|
|
||||||
ID: "tc_1",
|
|
||||||
Function: llm.FunctionCall{Name: "slow_work", Arguments: `{}`},
|
|
||||||
}),
|
|
||||||
// After resume: the incremental checkpoint saved after tool completion
|
|
||||||
// means restore continues with these messages; LLM returns final answer.
|
|
||||||
stopResponse("All done after resume."),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
ag := agent.New(
|
|
||||||
"worker-agent",
|
|
||||||
newTestClient(provider),
|
|
||||||
agent.WithModel("test-model"),
|
|
||||||
agent.WithTools(slowTool),
|
|
||||||
)
|
|
||||||
|
|
||||||
registry := &simpleRegistry{
|
|
||||||
agents: map[string]*agent.Agent{"worker-agent": ag},
|
|
||||||
}
|
|
||||||
|
|
||||||
run := agentruntest.InsertPendingRun(
|
|
||||||
t,
|
|
||||||
client,
|
|
||||||
"worker-agent",
|
|
||||||
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "do work"}}}},
|
|
||||||
)
|
|
||||||
|
|
||||||
supervisor := probo.NewAgentRunSupervisor(
|
|
||||||
client,
|
|
||||||
store,
|
|
||||||
registry,
|
|
||||||
testLogger(),
|
|
||||||
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
|
||||||
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
|
||||||
)
|
|
||||||
|
|
||||||
// --- Phase 1: Start and let the supervisor pick up the run ---
|
|
||||||
|
|
||||||
ctx1, cancel1 := context.WithTimeout(context.Background(), 15*time.Second)
|
|
||||||
defer cancel1()
|
|
||||||
|
|
||||||
go func() { _ = supervisor.Run(ctx1) }()
|
|
||||||
|
|
||||||
// Wait for the tool to start executing — this confirms the supervisor
|
|
||||||
// claimed the run and the agent called the tool.
|
|
||||||
select {
|
|
||||||
case <-toolReady:
|
|
||||||
case <-ctx1.Done():
|
|
||||||
t.Fatal("timed out waiting for tool to start")
|
|
||||||
}
|
|
||||||
|
|
||||||
// The run should now be RUNNING.
|
|
||||||
running := agentruntest.LoadAgentRun(t, client, run.ID)
|
|
||||||
assert.Equal(t, coredata.AgentRunStatusRunning, running.Status)
|
|
||||||
|
|
||||||
// Trigger graceful shutdown of the supervisor: context.AfterFunc
|
|
||||||
// registered in Run() fires signalShutdown, closing the shutdown
|
|
||||||
// broadcast channel; the per-run forwarder goroutine closes the
|
|
||||||
// agent's stop channel.
|
|
||||||
cancel1()
|
|
||||||
|
|
||||||
// Wait for the shutdown broadcast to be observed (the AfterFunc
|
|
||||||
// goroutine closes it) before releasing the tool. This is
|
|
||||||
// deterministic: no wall-clock sleep. The per-run forwarder
|
|
||||||
// goroutine observes the same close synchronously and closes the
|
|
||||||
// agent stop channel.
|
|
||||||
select {
|
|
||||||
case <-supervisor.ShutdownBroadcastForTests():
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("timed out waiting for supervisor shutdown broadcast")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Now release the tool. When the coreLoop resumes control at the
|
|
||||||
// next turn boundary it observes the closed stop channel, saves
|
|
||||||
// the suspension checkpoint, and returns SuspendedError.
|
|
||||||
close(toolRelease)
|
|
||||||
|
|
||||||
// Wait for the checkpoint to appear. The supervisor leaves the row
|
|
||||||
// in RUNNING because SuspendedError triggers the "leaving for stale
|
|
||||||
// recovery" path.
|
|
||||||
require.Eventually(
|
|
||||||
t,
|
|
||||||
func() bool {
|
|
||||||
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
|
||||||
return err == nil && r.Checkpoint != nil
|
|
||||||
},
|
|
||||||
10*time.Second,
|
|
||||||
200*time.Millisecond,
|
|
||||||
"checkpoint should be saved after stop",
|
|
||||||
)
|
|
||||||
|
|
||||||
// Verify checkpoint content.
|
|
||||||
cp, err := store.Load(context.Background(), run.ID.String())
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, cp, "checkpoint must exist after suspension")
|
|
||||||
assert.Equal(t, agent.AgentStatusSuspended, cp.Status)
|
|
||||||
assert.Equal(t, "worker-agent", cp.AgentName)
|
|
||||||
assert.True(t, len(cp.Messages) > 0, "checkpoint should contain messages")
|
|
||||||
|
|
||||||
// --- Phase 2: Simulate resume by resetting to PENDING ---
|
|
||||||
|
|
||||||
err = client.WithConn(
|
|
||||||
context.Background(),
|
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
_, err := conn.Exec(
|
|
||||||
ctx,
|
|
||||||
`UPDATE agent_runs
|
|
||||||
SET status = 'PENDING',
|
|
||||||
started_at = NULL,
|
|
||||||
lease_owner = NULL,
|
|
||||||
lease_expires_at = NULL,
|
|
||||||
updated_at = now()
|
|
||||||
WHERE id = $1`,
|
|
||||||
run.ID.String(),
|
|
||||||
)
|
|
||||||
|
|
||||||
return err
|
|
||||||
},
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Start a fresh supervisor to pick up the resumed run.
|
|
||||||
supervisor2 := probo.NewAgentRunSupervisor(
|
|
||||||
client,
|
|
||||||
store,
|
|
||||||
registry,
|
|
||||||
testLogger(),
|
|
||||||
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
|
||||||
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
|
||||||
)
|
|
||||||
|
|
||||||
ctx2, cancel2 := context.WithTimeout(context.Background(), 15*time.Second)
|
|
||||||
defer cancel2()
|
|
||||||
|
|
||||||
go func() { _ = supervisor2.Run(ctx2) }()
|
|
||||||
|
|
||||||
// The resumed run should load the checkpoint, call Restore, get the
|
|
||||||
// second LLM response (stopResponse), and complete.
|
|
||||||
require.Eventually(
|
|
||||||
t,
|
|
||||||
func() bool {
|
|
||||||
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
|
||||||
return err == nil && r.Status == coredata.AgentRunStatusCompleted
|
|
||||||
},
|
|
||||||
10*time.Second,
|
|
||||||
200*time.Millisecond,
|
|
||||||
"run should reach COMPLETED after resume",
|
|
||||||
)
|
|
||||||
|
|
||||||
completed := agentruntest.LoadAgentRun(t, client, run.ID)
|
|
||||||
assert.Equal(t, coredata.AgentRunStatusCompleted, completed.Status)
|
|
||||||
assert.NotNil(t, completed.Result)
|
|
||||||
assert.Nil(t, completed.Checkpoint, "checkpoint should be cleared after completion")
|
|
||||||
assert.Nil(t, completed.ErrorMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Test 5: SIGTERM battle test — realistic multi-turn security audit
|
|
||||||
// with parallel tool calls, long-running operations, thinking turns,
|
|
||||||
// and multiple kill/resume cycles.
|
|
||||||
//
|
|
||||||
// Simulated workflow (10 tool-call turns + 1 final response):
|
|
||||||
//
|
|
||||||
// Turn 0: [think] scan_repos (single, 800ms)
|
|
||||||
// Turn 1: [think] fetch_config ×3 (parallel, 300-500ms)
|
|
||||||
// Turn 2: [think] analyze (single, 1000ms)
|
|
||||||
// Turn 3: [think] check ×3 (parallel, 400-800ms)
|
|
||||||
// Turn 4: [think] deep_analysis (single, 1500ms — long running)
|
|
||||||
// Turn 5: [think] generate ×2 (parallel, 500-600ms)
|
|
||||||
// Turn 6: [think] cve_lookup (single, 700ms)
|
|
||||||
// Turn 7: [think] compile (single, 600ms)
|
|
||||||
// Turn 8: [think] validate + format (parallel, 300-400ms)
|
|
||||||
// Turn 9: [think] publish (single, 400ms)
|
|
||||||
// Turn 10: final response — "Security audit complete..."
|
|
||||||
//
|
|
||||||
// SIGTERM is sent 3 times at different points, each time interrupting
|
|
||||||
// during tool execution (sometimes single, sometimes parallel).
|
|
||||||
// After each kill the checkpoint is verified to show progressive
|
|
||||||
// accumulation. A final in-process resume runs the remaining turns
|
|
||||||
// to completion.
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
// workInput is the shared parameter type for all battle-test tools.
|
|
||||||
type workInput struct {
|
|
||||||
Task string `json:"task"`
|
|
||||||
DurationMs int `json:"duration_ms"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// battleTestResponses returns the full LLM response sequence for a
|
|
||||||
// simulated security-audit agent. Each tool-call turn includes
|
|
||||||
// thinking text so the checkpoint messages are realistic.
|
|
||||||
func battleTestResponses() []*llm.ChatCompletionResponse {
|
|
||||||
tc := func(id, name, task string, ms int) llm.ToolCall {
|
|
||||||
return llm.ToolCall{
|
|
||||||
ID: id,
|
|
||||||
Function: llm.FunctionCall{
|
|
||||||
Name: name,
|
|
||||||
Arguments: fmt.Sprintf(`{"task":%q,"duration_ms":%d}`, task, ms),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
think := func(text string, calls ...llm.ToolCall) *llm.ChatCompletionResponse {
|
|
||||||
return &llm.ChatCompletionResponse{
|
|
||||||
Model: "test-model",
|
|
||||||
Message: llm.Message{
|
|
||||||
Role: llm.RoleAssistant,
|
|
||||||
Parts: []llm.Part{llm.TextPart{Text: text}},
|
|
||||||
ToolCalls: calls,
|
|
||||||
},
|
|
||||||
Usage: llm.Usage{InputTokens: 50, OutputTokens: 30},
|
|
||||||
FinishReason: llm.FinishReasonToolCalls,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return []*llm.ChatCompletionResponse{
|
|
||||||
// Turn 0 — single long scan
|
|
||||||
think(
|
|
||||||
"I'll begin the security audit by scanning all repositories to identify codebases, dependency manifests, and access-control configurations.",
|
|
||||||
tc("tc_0_1", "scan", "scan_repos", 800),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 1 — 3 parallel fetches
|
|
||||||
think(
|
|
||||||
"Found 3 repositories: api-gateway, auth-service, data-pipeline. Fetching their configurations in parallel to save time.",
|
|
||||||
tc("tc_1_1", "fetch", "fetch_api_config", 300),
|
|
||||||
tc("tc_1_2", "fetch", "fetch_auth_config", 400),
|
|
||||||
tc("tc_1_3", "fetch", "fetch_data_config", 500),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 2 — single analysis
|
|
||||||
think(
|
|
||||||
"All configurations retrieved. Running a comprehensive vulnerability analysis against the OWASP Top-10 checklist.",
|
|
||||||
tc("tc_2_1", "analyze", "analyze_configs", 1000),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 3 — 3 parallel security checks
|
|
||||||
think(
|
|
||||||
"Analysis flagged several areas of concern. Running dependency audit, secret scanning, and IAM permission checks in parallel.",
|
|
||||||
tc("tc_3_1", "check", "check_dependencies", 600),
|
|
||||||
tc("tc_3_2", "check", "check_secrets", 400),
|
|
||||||
tc("tc_3_3", "check", "check_permissions", 800),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 4 — single very long deep-dive
|
|
||||||
think(
|
|
||||||
"Multiple issues found: 3 outdated dependencies with known CVEs, 2 overly permissive IAM roles. Performing a deep analysis on the critical findings to determine exploitability and blast radius.",
|
|
||||||
tc("tc_4_1", "analyze", "deep_analysis", 1500),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 5 — 2 parallel report sections
|
|
||||||
think(
|
|
||||||
"Deep analysis complete. auth-service uses deprecated TLS 1.1 and data-pipeline stores PII unencrypted. Generating the executive summary and detailed findings sections in parallel.",
|
|
||||||
tc("tc_5_1", "generate", "generate_summary", 500),
|
|
||||||
tc("tc_5_2", "generate", "generate_findings", 600),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 6 — single CVE lookup
|
|
||||||
think(
|
|
||||||
"Report sections drafted. Cross-referencing all findings against the NVD and GitHub Advisory databases for known CVE identifiers.",
|
|
||||||
tc("tc_6_1", "lookup", "cve_lookup", 700),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 7 — single compile
|
|
||||||
think(
|
|
||||||
"CVE-2026-1234 matches the auth-service TLS vulnerability (CVSS 9.1). Compiling all sections, references, and remediation steps into the final report.",
|
|
||||||
tc("tc_7_1", "compile", "compile_report", 600),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 8 — 2 parallel validation + formatting
|
|
||||||
think(
|
|
||||||
"Draft report assembled (12 pages). Running structural validation and PDF formatting concurrently.",
|
|
||||||
tc("tc_8_1", "validate", "validate_report", 400),
|
|
||||||
tc("tc_8_2", "format", "format_pdf", 300),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 9 — single publish
|
|
||||||
think(
|
|
||||||
"Validation passed, PDF formatted. Publishing the finalized audit report to the internal portal.",
|
|
||||||
tc("tc_9_1", "publish", "publish_report", 400),
|
|
||||||
),
|
|
||||||
|
|
||||||
// Turn 10 — final text response
|
|
||||||
{
|
|
||||||
Model: "test-model",
|
|
||||||
Message: llm.Message{
|
|
||||||
Role: llm.RoleAssistant,
|
|
||||||
Parts: []llm.Part{llm.TextPart{Text: "Security audit complete.\n\nFindings:\n- 3 critical (auth-service TLS 1.1, unencrypted PII, CVE-2026-1234)\n- 5 medium (outdated deps, permissive IAM)\n- 4 low (missing rate-limiting, verbose logging)\n\nFull report: https://audits.internal/report-2026-04"}},
|
|
||||||
},
|
|
||||||
Usage: llm.Usage{InputTokens: 50, OutputTokens: 30},
|
|
||||||
FinishReason: llm.FinishReasonStop,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// makeBattleTools creates the tool set for the battle test. Every tool
|
|
||||||
// shares the same handler that sleeps for the requested duration and
|
|
||||||
// records progress to a shared file.
|
|
||||||
func makeBattleTools(progressFile string) []agent.Tool {
|
|
||||||
var mu sync.Mutex
|
|
||||||
|
|
||||||
handler := func(_ context.Context, input workInput) (agent.ToolResult, error) {
|
|
||||||
// Simulate real work.
|
|
||||||
time.Sleep(time.Duration(input.DurationMs) * time.Millisecond)
|
|
||||||
|
|
||||||
// Record completion — written AFTER the sleep so the parent's
|
|
||||||
// step count reflects truly-finished work.
|
|
||||||
mu.Lock()
|
|
||||||
|
|
||||||
f, err := os.OpenFile(progressFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
|
|
||||||
if err != nil {
|
|
||||||
mu.Unlock()
|
|
||||||
return agent.ToolResult{}, err
|
|
||||||
}
|
|
||||||
|
|
||||||
_, _ = fmt.Fprintln(f, input.Task)
|
|
||||||
_ = f.Close()
|
|
||||||
mu.Unlock()
|
|
||||||
|
|
||||||
return agent.ToolResult{Content: fmt.Sprintf("completed: %s", input.Task)}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
names := []struct{ name, desc string }{
|
|
||||||
{"scan", "Scan repositories for audit targets"},
|
|
||||||
{"fetch", "Fetch configuration or source files"},
|
|
||||||
{"analyze", "Run vulnerability analysis"},
|
|
||||||
{"check", "Execute a specific security check"},
|
|
||||||
{"generate", "Generate a report section"},
|
|
||||||
{"lookup", "Query external vulnerability databases"},
|
|
||||||
{"compile", "Compile report sections into final document"},
|
|
||||||
{"validate", "Validate report structure"},
|
|
||||||
{"format", "Apply output formatting"},
|
|
||||||
{"publish", "Publish report to internal portal"},
|
|
||||||
}
|
|
||||||
|
|
||||||
tools := make([]agent.Tool, len(names))
|
|
||||||
for i, n := range names {
|
|
||||||
tools[i] = agent.FunctionTool[workInput](n.name, n.desc, handler)
|
|
||||||
}
|
|
||||||
|
|
||||||
return tools
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAgentRunSupervisor_SIGTERM(t *testing.T) {
|
|
||||||
// ---- Subprocess mode ----
|
|
||||||
if os.Getenv("TEST_SIGTERM_SUBPROCESS") == "1" {
|
|
||||||
runSIGTERMSubprocess()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- Parent mode ----
|
|
||||||
client := agentruntest.PGClient(t)
|
|
||||||
store := coredata.NewPGCheckpointer(client)
|
|
||||||
|
|
||||||
run := agentruntest.InsertPendingRun(
|
|
||||||
t,
|
|
||||||
client,
|
|
||||||
"battle-agent",
|
|
||||||
[]llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Run a full security audit on all repositories."}}}},
|
|
||||||
)
|
|
||||||
|
|
||||||
progressFile := filepath.Join(t.TempDir(), "progress")
|
|
||||||
|
|
||||||
// ---- Helpers ----
|
|
||||||
|
|
||||||
countSteps := func() int {
|
|
||||||
data, err := os.ReadFile(progressFile)
|
|
||||||
if err != nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
n := 0
|
|
||||||
|
|
||||||
for _, b := range data {
|
|
||||||
if b == '\n' {
|
|
||||||
n++
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return n
|
|
||||||
}
|
|
||||||
|
|
||||||
startSubprocess := func(skipResponses int) *exec.Cmd {
|
|
||||||
cmd := exec.Command(
|
|
||||||
os.Args[0],
|
|
||||||
"-test.run=^TestAgentRunSupervisor_SIGTERM$",
|
|
||||||
"-test.v",
|
|
||||||
)
|
|
||||||
|
|
||||||
cmd.Env = append(os.Environ(),
|
|
||||||
"TEST_SIGTERM_SUBPROCESS=1",
|
|
||||||
"TEST_SIGTERM_PROGRESS_FILE="+progressFile,
|
|
||||||
"TEST_SIGTERM_SKIP_RESPONSES="+strconv.Itoa(skipResponses),
|
|
||||||
)
|
|
||||||
cmd.Stdout = os.Stdout
|
|
||||||
cmd.Stderr = os.Stderr
|
|
||||||
require.NoError(t, cmd.Start())
|
|
||||||
|
|
||||||
return cmd
|
|
||||||
}
|
|
||||||
|
|
||||||
killAndWait := func(cmd *exec.Cmd) {
|
|
||||||
require.NoError(t, cmd.Process.Signal(syscall.SIGTERM))
|
|
||||||
|
|
||||||
err := cmd.Wait()
|
|
||||||
if err == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
exitErr, ok := errors.AsType[*exec.ExitError](err)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("subprocess error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ws, ok := exitErr.Sys().(syscall.WaitStatus)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("subprocess exited with unexpected wait status: %v", exitErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
if ws.Signaled() && ws.Signal() == syscall.SIGTERM {
|
|
||||||
t.Logf("subprocess terminated by SIGTERM")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Fatalf("subprocess exited unexpectedly (signaled=%v signal=%v exit=%d): %v",
|
|
||||||
ws.Signaled(), ws.Signal(), ws.ExitStatus(), exitErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
resetToPending := func() {
|
|
||||||
err := client.WithConn(
|
|
||||||
context.Background(),
|
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
_, err := conn.Exec(ctx, `
|
|
||||||
UPDATE agent_runs
|
|
||||||
SET status = 'PENDING',
|
|
||||||
started_at = NULL,
|
|
||||||
lease_owner = NULL,
|
|
||||||
lease_expires_at = NULL,
|
|
||||||
updated_at = now()
|
|
||||||
WHERE id = $1`,
|
|
||||||
run.ID.String(),
|
|
||||||
)
|
|
||||||
|
|
||||||
return err
|
|
||||||
},
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
waitForSteps := func(target int) {
|
|
||||||
require.Eventually(
|
|
||||||
t,
|
|
||||||
func() bool { return countSteps() >= target },
|
|
||||||
30*time.Second,
|
|
||||||
100*time.Millisecond,
|
|
||||||
fmt.Sprintf("expected at least %d completed tool executions", target),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
verifyCheckpoint := func(phase int) *agent.Checkpoint {
|
|
||||||
cp, err := store.Load(context.Background(), run.ID.String())
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, cp, "phase %d: checkpoint must exist", phase)
|
|
||||||
assert.Equal(t, agent.AgentStatusSuspended, cp.Status, "phase %d", phase)
|
|
||||||
assert.Equal(t, "battle-agent", cp.AgentName, "phase %d", phase)
|
|
||||||
assert.Greater(t, len(cp.Messages), 1, "phase %d: checkpoint should have messages", phase)
|
|
||||||
assert.Greater(t, cp.Turns, 0, "phase %d: checkpoint should have turns", phase)
|
|
||||||
t.Logf(
|
|
||||||
" checkpoint: %d messages, %d turns, usage=%+v",
|
|
||||||
len(cp.Messages), cp.Turns, cp.Usage,
|
|
||||||
)
|
|
||||||
|
|
||||||
return cp
|
|
||||||
}
|
|
||||||
|
|
||||||
// ============================================================
|
|
||||||
// Phase 1: SIGTERM during the parallel fetch (turn 1)
|
|
||||||
// Steps so far: turn0=1(scan) + turn1=3(fetch×3) = 4
|
|
||||||
// ============================================================
|
|
||||||
t.Log("=== Phase 1: SIGTERM after scan + parallel fetch (4 steps) ===")
|
|
||||||
|
|
||||||
cmd1 := startSubprocess(0)
|
|
||||||
|
|
||||||
waitForSteps(4)
|
|
||||||
killAndWait(cmd1)
|
|
||||||
|
|
||||||
steps1 := countSteps()
|
|
||||||
t.Logf(" %d tool executions completed", steps1)
|
|
||||||
require.GreaterOrEqual(t, steps1, 4)
|
|
||||||
|
|
||||||
cp1 := verifyCheckpoint(1)
|
|
||||||
require.GreaterOrEqual(t, cp1.Turns, 2, "should have completed at least turns 0-1")
|
|
||||||
|
|
||||||
resetToPending()
|
|
||||||
|
|
||||||
// ============================================================
|
|
||||||
// Phase 2: SIGTERM during the parallel security checks (turn 3)
|
|
||||||
// New steps: turn2=1(analyze) + turn3=3(check×3) = 4
|
|
||||||
// ============================================================
|
|
||||||
t.Log("=== Phase 2: SIGTERM after analyze + parallel checks (4 more steps) ===")
|
|
||||||
|
|
||||||
cmd2 := startSubprocess(cp1.Turns)
|
|
||||||
|
|
||||||
waitForSteps(steps1 + 4)
|
|
||||||
killAndWait(cmd2)
|
|
||||||
|
|
||||||
steps2 := countSteps()
|
|
||||||
t.Logf(" %d tool executions completed (total)", steps2)
|
|
||||||
require.GreaterOrEqual(t, steps2, steps1+4)
|
|
||||||
|
|
||||||
cp2 := verifyCheckpoint(2)
|
|
||||||
assert.Greater(t, cp2.Turns, cp1.Turns, "turns should grow")
|
|
||||||
assert.Greater(t, len(cp2.Messages), len(cp1.Messages), "messages should grow")
|
|
||||||
|
|
||||||
resetToPending()
|
|
||||||
|
|
||||||
// ============================================================
|
|
||||||
// Phase 3: SIGTERM during deep_analysis (turn 4, long-running)
|
|
||||||
// or after generate ×2 (turn 5)
|
|
||||||
// New steps: turn4=1(deep) + turn5=2(generate×2) = 3
|
|
||||||
// ============================================================
|
|
||||||
t.Log("=== Phase 3: SIGTERM during long-running deep analysis (3 more steps) ===")
|
|
||||||
|
|
||||||
cmd3 := startSubprocess(cp2.Turns)
|
|
||||||
|
|
||||||
waitForSteps(steps2 + 3)
|
|
||||||
killAndWait(cmd3)
|
|
||||||
|
|
||||||
steps3 := countSteps()
|
|
||||||
t.Logf(" %d tool executions completed (total)", steps3)
|
|
||||||
require.GreaterOrEqual(t, steps3, steps2+3)
|
|
||||||
|
|
||||||
cp3 := verifyCheckpoint(3)
|
|
||||||
assert.Greater(t, cp3.Turns, cp2.Turns, "turns should grow again")
|
|
||||||
assert.Greater(t, len(cp3.Messages), len(cp2.Messages), "messages should grow again")
|
|
||||||
assert.Greater(t, cp3.Usage.InputTokens, 0, "usage should accumulate")
|
|
||||||
assert.Greater(t, cp3.Usage.OutputTokens, 0, "usage should accumulate")
|
|
||||||
|
|
||||||
t.Logf(
|
|
||||||
" after 3 SIGTERM cycles: %d steps, %d turns, %d messages, usage=%+v",
|
|
||||||
steps3, cp3.Turns, len(cp3.Messages), cp3.Usage,
|
|
||||||
)
|
|
||||||
|
|
||||||
resetToPending()
|
|
||||||
|
|
||||||
// ============================================================
|
|
||||||
// Phase 4: final in-process resume — run remaining turns to
|
|
||||||
// completion (lookup, compile, validate+format, publish, done)
|
|
||||||
// ============================================================
|
|
||||||
t.Log("=== Phase 4: in-process resume to completion ===")
|
|
||||||
|
|
||||||
remaining := battleTestResponses()[cp3.Turns:]
|
|
||||||
t.Logf(" %d LLM responses remaining (turns %d–10)", len(remaining), cp3.Turns)
|
|
||||||
|
|
||||||
tools := makeBattleTools(progressFile)
|
|
||||||
|
|
||||||
resumeAgent := agent.New(
|
|
||||||
"battle-agent",
|
|
||||||
newTestClient(&mockProvider{responses: remaining}),
|
|
||||||
agent.WithModel("test-model"),
|
|
||||||
agent.WithTools(tools...),
|
|
||||||
agent.WithMaxTurns(25),
|
|
||||||
)
|
|
||||||
|
|
||||||
supervisor := probo.NewAgentRunSupervisor(
|
|
||||||
client,
|
|
||||||
store,
|
|
||||||
&simpleRegistry{agents: map[string]*agent.Agent{"battle-agent": resumeAgent}},
|
|
||||||
testLogger(),
|
|
||||||
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
|
||||||
probo.WithAgentRunSupervisorLeaseDuration(30*time.Second),
|
|
||||||
)
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
go func() { _ = supervisor.Run(ctx) }()
|
|
||||||
|
|
||||||
require.Eventually(
|
|
||||||
t,
|
|
||||||
func() bool {
|
|
||||||
r, err := agentruntest.TryLoadAgentRun(client, run.ID)
|
|
||||||
return err == nil && r.Status == coredata.AgentRunStatusCompleted
|
|
||||||
},
|
|
||||||
25*time.Second,
|
|
||||||
200*time.Millisecond,
|
|
||||||
"run should complete after final resume",
|
|
||||||
)
|
|
||||||
|
|
||||||
stepsFinal := countSteps()
|
|
||||||
final := agentruntest.LoadAgentRun(t, client, run.ID)
|
|
||||||
assert.Equal(t, coredata.AgentRunStatusCompleted, final.Status)
|
|
||||||
assert.NotNil(t, final.Result)
|
|
||||||
assert.Nil(t, final.Checkpoint, "checkpoint should be cleared")
|
|
||||||
assert.Nil(t, final.ErrorMessage)
|
|
||||||
assert.Contains(t, string(final.Result), "Security audit complete")
|
|
||||||
|
|
||||||
t.Logf(
|
|
||||||
" battle test done: %d total tool executions across 3 SIGTERM cycles + final resume",
|
|
||||||
stepsFinal,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// runSIGTERMSubprocess is the child-process entry point. It sets up a
|
|
||||||
// supervisor with the full security-audit agent (10 distinct tools,
|
|
||||||
// thinking text, parallel calls, varying durations) and handles
|
|
||||||
// SIGTERM via signal.NotifyContext — identical to production probod.
|
|
||||||
func runSIGTERMSubprocess() {
|
|
||||||
progressFile := os.Getenv("TEST_SIGTERM_PROGRESS_FILE")
|
|
||||||
skip, _ := strconv.Atoi(os.Getenv("TEST_SIGTERM_SKIP_RESPONSES"))
|
|
||||||
|
|
||||||
addr := os.Getenv("PROBO_TEST_PG_ADDR")
|
|
||||||
if addr == "" {
|
|
||||||
addr = "localhost:5432"
|
|
||||||
}
|
|
||||||
|
|
||||||
user := os.Getenv("PROBO_TEST_PG_USER")
|
|
||||||
if user == "" {
|
|
||||||
user = "probod"
|
|
||||||
}
|
|
||||||
|
|
||||||
password := os.Getenv("PROBO_TEST_PG_PASSWORD")
|
|
||||||
if password == "" {
|
|
||||||
password = "probod"
|
|
||||||
}
|
|
||||||
|
|
||||||
database := os.Getenv("PROBO_TEST_PG_DATABASE")
|
|
||||||
if database == "" {
|
|
||||||
database = "probod_test"
|
|
||||||
}
|
|
||||||
|
|
||||||
pgClient, err := pg.NewClient(
|
|
||||||
pg.WithAddr(addr),
|
|
||||||
pg.WithUser(user),
|
|
||||||
pg.WithPassword(password),
|
|
||||||
pg.WithDatabase(database),
|
|
||||||
pg.WithPoolSize(5),
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
fmt.Fprintf(os.Stderr, "subprocess: cannot create pg client: %v\n", err)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
defer pgClient.Close()
|
|
||||||
|
|
||||||
store := coredata.NewPGCheckpointer(pgClient)
|
|
||||||
tools := makeBattleTools(progressFile)
|
|
||||||
|
|
||||||
responses := battleTestResponses()
|
|
||||||
if skip > 0 && skip < len(responses) {
|
|
||||||
responses = responses[skip:]
|
|
||||||
}
|
|
||||||
|
|
||||||
ag := agent.New(
|
|
||||||
"battle-agent",
|
|
||||||
newTestClient(&mockProvider{responses: responses}),
|
|
||||||
agent.WithModel("test-model"),
|
|
||||||
agent.WithTools(tools...),
|
|
||||||
agent.WithMaxTurns(25),
|
|
||||||
)
|
|
||||||
|
|
||||||
supervisor := probo.NewAgentRunSupervisor(
|
|
||||||
pgClient,
|
|
||||||
store,
|
|
||||||
&simpleRegistry{agents: map[string]*agent.Agent{"battle-agent": ag}},
|
|
||||||
log.NewLogger(log.WithFormat(log.FormatPretty)),
|
|
||||||
probo.WithAgentRunSupervisorInterval(500*time.Millisecond),
|
|
||||||
probo.WithAgentRunSupervisorLeaseDuration(5*time.Second),
|
|
||||||
)
|
|
||||||
|
|
||||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGTERM)
|
|
||||||
defer stop()
|
|
||||||
|
|
||||||
// signal.NotifyContext sets the cancellation cause to the signal,
|
|
||||||
// so context.Cause(ctx) returns syscall.SIGTERM rather than
|
|
||||||
// context.Canceled. Treat any ctx-cancellation outcome as graceful.
|
|
||||||
err = supervisor.Run(ctx)
|
|
||||||
if err != nil && ctx.Err() == nil {
|
|
||||||
fmt.Fprintf(os.Stderr, "subprocess: supervisor error: %v\n", err)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
os.Exit(0)
|
|
||||||
}
|
|
||||||
@@ -1,247 +0,0 @@
|
|||||||
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
|
||||||
//
|
|
||||||
// Permission to use, copy, modify, and/or distribute this software for any
|
|
||||||
// purpose with or without fee is hereby granted, provided that the above
|
|
||||||
// copyright notice and this permission notice appear in all copies.
|
|
||||||
//
|
|
||||||
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
|
||||||
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
|
||||||
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
|
||||||
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
|
||||||
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
|
||||||
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
|
||||||
// PERFORMANCE OF THIS SOFTWARE.
|
|
||||||
|
|
||||||
// Package agentruntest provides shared test helpers for agent run
|
|
||||||
// integration tests that require a PostgreSQL database.
|
|
||||||
package agentruntest
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"go.gearno.de/kit/pg"
|
|
||||||
"go.probo.inc/probo/pkg/coredata"
|
|
||||||
"go.probo.inc/probo/pkg/gid"
|
|
||||||
"go.probo.inc/probo/pkg/llm"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
sharedPGClient *pg.Client
|
|
||||||
pgOnce sync.Once
|
|
||||||
pgInitErr error
|
|
||||||
ensureTableOnce sync.Once
|
|
||||||
ensureTableErr error
|
|
||||||
)
|
|
||||||
|
|
||||||
// PGClient returns a shared pg.Client connected to the test database.
|
|
||||||
// Skips the test if the database is not reachable.
|
|
||||||
func PGClient(t *testing.T) *pg.Client {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
pgOnce.Do(func() {
|
|
||||||
addr := os.Getenv("PROBO_TEST_PG_ADDR")
|
|
||||||
if addr == "" {
|
|
||||||
addr = "localhost:5432"
|
|
||||||
}
|
|
||||||
|
|
||||||
user := os.Getenv("PROBO_TEST_PG_USER")
|
|
||||||
if user == "" {
|
|
||||||
user = "probod"
|
|
||||||
}
|
|
||||||
|
|
||||||
password := os.Getenv("PROBO_TEST_PG_PASSWORD")
|
|
||||||
if password == "" {
|
|
||||||
password = "probod"
|
|
||||||
}
|
|
||||||
|
|
||||||
database := os.Getenv("PROBO_TEST_PG_DATABASE")
|
|
||||||
if database == "" {
|
|
||||||
database = "probod_test"
|
|
||||||
}
|
|
||||||
|
|
||||||
sharedPGClient, pgInitErr = pg.NewClient(
|
|
||||||
pg.WithAddr(addr),
|
|
||||||
pg.WithUser(user),
|
|
||||||
pg.WithPassword(password),
|
|
||||||
pg.WithDatabase(database),
|
|
||||||
pg.WithPoolSize(5),
|
|
||||||
)
|
|
||||||
if pgInitErr != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
pgInitErr = sharedPGClient.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
_, err := conn.Exec(ctx, "SELECT 1")
|
|
||||||
return err
|
|
||||||
})
|
|
||||||
})
|
|
||||||
|
|
||||||
if pgInitErr != nil {
|
|
||||||
t.Skipf("cannot connect to test database: %v", pgInitErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
EnsureAgentRunsTable(t, sharedPGClient)
|
|
||||||
|
|
||||||
return sharedPGClient
|
|
||||||
}
|
|
||||||
|
|
||||||
// EnsureAgentRunsTable creates the agent_runs table against the test
|
|
||||||
// database using the embedded migration, if the table is not already
|
|
||||||
// present. If the table exists with a stale schema (e.g. missing the
|
|
||||||
// FK added later), drop it manually or let the production migration
|
|
||||||
// runner apply the current version — this helper does not rewrite an
|
|
||||||
// existing table to avoid racing concurrent test processes that share
|
|
||||||
// the same database.
|
|
||||||
func EnsureAgentRunsTable(t *testing.T, client *pg.Client) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ensureTableOnce.Do(func() {
|
|
||||||
ctx := context.Background()
|
|
||||||
ensureTableErr = client.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
var exists bool
|
|
||||||
if err := conn.QueryRow(
|
|
||||||
ctx,
|
|
||||||
`SELECT EXISTS (SELECT 1 FROM information_schema.tables WHERE table_name = 'agent_runs')`,
|
|
||||||
).Scan(&exists); err != nil {
|
|
||||||
return fmt.Errorf("cannot check agent_runs existence: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if exists {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
ddl, err := coredata.Migrations.ReadFile("migrations/20260424T173529Z.sql")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("cannot read agent_runs migration: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := conn.Exec(ctx, string(ddl)); err != nil {
|
|
||||||
return fmt.Errorf("cannot apply agent_runs migration: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
})
|
|
||||||
require.NoError(t, ensureTableErr, "cannot ensure agent_runs table")
|
|
||||||
}
|
|
||||||
|
|
||||||
// CleanupAgentRun deletes an agent run by ID. Safe to call from
|
|
||||||
// t.Cleanup.
|
|
||||||
func CleanupAgentRun(client *pg.Client, id gid.GID) {
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
_ = client.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
_, err := conn.Exec(ctx, "DELETE FROM agent_runs WHERE id = $1", id.String())
|
|
||||||
return err
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// InsertPendingRun inserts a PENDING agent run and registers cleanup.
|
|
||||||
// A placeholder organization row is created first so the agent_runs FK
|
|
||||||
// on organization_id is satisfied.
|
|
||||||
func InsertPendingRun(
|
|
||||||
t *testing.T,
|
|
||||||
client *pg.Client,
|
|
||||||
agentName string,
|
|
||||||
inputMessages []llm.Message,
|
|
||||||
) coredata.AgentRun {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
tenantID := gid.NewTenantID()
|
|
||||||
orgID := gid.New(tenantID, 1)
|
|
||||||
runID := gid.New(tenantID, 2)
|
|
||||||
|
|
||||||
inputJSON, err := json.Marshal(inputMessages)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
now := time.Now()
|
|
||||||
|
|
||||||
run := coredata.AgentRun{
|
|
||||||
ID: runID,
|
|
||||||
OrganizationID: orgID,
|
|
||||||
StartAgentName: agentName,
|
|
||||||
Status: coredata.AgentRunStatusPending,
|
|
||||||
InputMessages: inputJSON,
|
|
||||||
CreatedAt: now,
|
|
||||||
UpdatedAt: now,
|
|
||||||
}
|
|
||||||
|
|
||||||
err = client.WithTx(
|
|
||||||
context.Background(),
|
|
||||||
func(ctx context.Context, tx pg.Tx) error {
|
|
||||||
if _, err := tx.Exec(
|
|
||||||
ctx,
|
|
||||||
`INSERT INTO organizations (id, tenant_id, name, created_at, updated_at) VALUES ($1, $2, $3, $4, $5)`,
|
|
||||||
orgID.String(), tenantID.String(), "test-org-"+orgID.String(), now, now,
|
|
||||||
); err != nil {
|
|
||||||
return fmt.Errorf("cannot insert placeholder organization: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return run.Insert(ctx, tx, coredata.NewScope(tenantID))
|
|
||||||
},
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
cleanupOrganization(client, orgID)
|
|
||||||
})
|
|
||||||
|
|
||||||
return run
|
|
||||||
}
|
|
||||||
|
|
||||||
// cleanupOrganization deletes the test organization row; the agent_runs
|
|
||||||
// FK has ON DELETE CASCADE so the associated run is removed too.
|
|
||||||
func cleanupOrganization(client *pg.Client, id gid.GID) {
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
_ = client.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
_, err := conn.Exec(ctx, "DELETE FROM organizations WHERE id = $1", id.String())
|
|
||||||
return err
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// LoadAgentRun loads an agent run by ID, failing the test on error.
|
|
||||||
func LoadAgentRun(t *testing.T, client *pg.Client, id gid.GID) coredata.AgentRun {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var run coredata.AgentRun
|
|
||||||
|
|
||||||
err := client.WithConn(
|
|
||||||
context.Background(),
|
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
return run.LoadByID(ctx, conn, coredata.NewNoScope(), id)
|
|
||||||
},
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("cannot load agent run %s: %v", id, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return run
|
|
||||||
}
|
|
||||||
|
|
||||||
// TryLoadAgentRun is a non-fatal variant safe for use inside
|
|
||||||
// require.Eventually callbacks (which recover panics).
|
|
||||||
func TryLoadAgentRun(client *pg.Client, id gid.GID) (coredata.AgentRun, error) {
|
|
||||||
var run coredata.AgentRun
|
|
||||||
|
|
||||||
err := client.WithConn(
|
|
||||||
context.Background(),
|
|
||||||
func(ctx context.Context, conn pg.Querier) error {
|
|
||||||
return run.LoadByID(ctx, conn, coredata.NewNoScope(), id)
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
return run, err
|
|
||||||
}
|
|
||||||
@@ -35,19 +35,19 @@ type (
|
|||||||
AgentRunStatus string
|
AgentRunStatus string
|
||||||
|
|
||||||
AgentRun struct {
|
AgentRun struct {
|
||||||
ID gid.GID `db:"id"`
|
ID gid.GID `db:"id"`
|
||||||
OrganizationID gid.GID `db:"organization_id"`
|
OrganizationID gid.GID `db:"organization_id"`
|
||||||
StartAgentName string `db:"start_agent_name"`
|
StartAgentName string `db:"start_agent_name"`
|
||||||
Status AgentRunStatus `db:"status"`
|
Status AgentRunStatus `db:"status"`
|
||||||
Checkpoint json.RawMessage `db:"checkpoint"`
|
Checkpoint json.RawMessage `db:"checkpoint"`
|
||||||
InputMessages json.RawMessage `db:"input_messages"`
|
InputMessages json.RawMessage `db:"input_messages"`
|
||||||
Result json.RawMessage `db:"result"`
|
Result json.RawMessage `db:"result"`
|
||||||
ErrorMessage *string `db:"error_message"`
|
ErrorMessage *string `db:"error_message"`
|
||||||
StartedAt *time.Time `db:"started_at"`
|
StartedAt *time.Time `db:"started_at"`
|
||||||
LeaseOwner *string `db:"lease_owner"`
|
LeaseExpiresAt *time.Time `db:"lease_expires_at"`
|
||||||
LeaseExpiresAt *time.Time `db:"lease_expires_at"`
|
LeaseGeneration int64 `db:"lease_generation"`
|
||||||
CreatedAt time.Time `db:"created_at"`
|
CreatedAt time.Time `db:"created_at"`
|
||||||
UpdatedAt time.Time `db:"updated_at"`
|
UpdatedAt time.Time `db:"updated_at"`
|
||||||
}
|
}
|
||||||
|
|
||||||
AgentRuns []*AgentRun
|
AgentRuns []*AgentRun
|
||||||
@@ -178,8 +178,8 @@ SELECT
|
|||||||
result,
|
result,
|
||||||
error_message,
|
error_message,
|
||||||
started_at,
|
started_at,
|
||||||
lease_owner,
|
|
||||||
lease_expires_at,
|
lease_expires_at,
|
||||||
|
lease_generation,
|
||||||
created_at,
|
created_at,
|
||||||
updated_at
|
updated_at
|
||||||
FROM
|
FROM
|
||||||
@@ -231,8 +231,8 @@ SELECT
|
|||||||
result,
|
result,
|
||||||
error_message,
|
error_message,
|
||||||
started_at,
|
started_at,
|
||||||
lease_owner,
|
|
||||||
lease_expires_at,
|
lease_expires_at,
|
||||||
|
lease_generation,
|
||||||
created_at,
|
created_at,
|
||||||
updated_at
|
updated_at
|
||||||
FROM
|
FROM
|
||||||
@@ -286,8 +286,8 @@ SELECT
|
|||||||
result,
|
result,
|
||||||
error_message,
|
error_message,
|
||||||
started_at,
|
started_at,
|
||||||
lease_owner,
|
|
||||||
lease_expires_at,
|
lease_expires_at,
|
||||||
|
lease_generation,
|
||||||
created_at,
|
created_at,
|
||||||
updated_at
|
updated_at
|
||||||
FROM
|
FROM
|
||||||
@@ -383,8 +383,8 @@ RETURNING
|
|||||||
result,
|
result,
|
||||||
error_message,
|
error_message,
|
||||||
started_at,
|
started_at,
|
||||||
lease_owner,
|
|
||||||
lease_expires_at,
|
lease_expires_at,
|
||||||
|
lease_generation,
|
||||||
created_at,
|
created_at,
|
||||||
updated_at;
|
updated_at;
|
||||||
`
|
`
|
||||||
@@ -432,8 +432,8 @@ SET
|
|||||||
result = @result,
|
result = @result,
|
||||||
error_message = @error_message,
|
error_message = @error_message,
|
||||||
started_at = @started_at,
|
started_at = @started_at,
|
||||||
lease_owner = @lease_owner,
|
|
||||||
lease_expires_at = @lease_expires_at,
|
lease_expires_at = @lease_expires_at,
|
||||||
|
lease_generation = @lease_generation,
|
||||||
updated_at = @updated_at
|
updated_at = @updated_at
|
||||||
WHERE
|
WHERE
|
||||||
%s
|
%s
|
||||||
@@ -448,8 +448,8 @@ RETURNING
|
|||||||
result,
|
result,
|
||||||
error_message,
|
error_message,
|
||||||
started_at,
|
started_at,
|
||||||
lease_owner,
|
|
||||||
lease_expires_at,
|
lease_expires_at,
|
||||||
|
lease_generation,
|
||||||
created_at,
|
created_at,
|
||||||
updated_at;
|
updated_at;
|
||||||
`
|
`
|
||||||
@@ -462,8 +462,8 @@ RETURNING
|
|||||||
"result": e.Result,
|
"result": e.Result,
|
||||||
"error_message": e.ErrorMessage,
|
"error_message": e.ErrorMessage,
|
||||||
"started_at": e.StartedAt,
|
"started_at": e.StartedAt,
|
||||||
"lease_owner": e.LeaseOwner,
|
|
||||||
"lease_expires_at": e.LeaseExpiresAt,
|
"lease_expires_at": e.LeaseExpiresAt,
|
||||||
|
"lease_generation": e.LeaseGeneration,
|
||||||
"updated_at": e.UpdatedAt,
|
"updated_at": e.UpdatedAt,
|
||||||
}
|
}
|
||||||
maps.Copy(args, scope.SQLArguments())
|
maps.Copy(args, scope.SQLArguments())
|
||||||
@@ -516,6 +516,46 @@ WHERE
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func CommitAgentRunResult(
|
||||||
|
ctx context.Context,
|
||||||
|
tx pg.Tx,
|
||||||
|
e *AgentRun,
|
||||||
|
leaseGeneration int64,
|
||||||
|
) (int64, error) {
|
||||||
|
q := `
|
||||||
|
UPDATE agent_runs
|
||||||
|
SET
|
||||||
|
status = @status,
|
||||||
|
result = @result,
|
||||||
|
error_message = @error_message,
|
||||||
|
started_at = @started_at,
|
||||||
|
lease_expires_at = @lease_expires_at,
|
||||||
|
updated_at = @updated_at
|
||||||
|
WHERE
|
||||||
|
id = @id
|
||||||
|
AND status = 'RUNNING'
|
||||||
|
AND lease_generation = @lease_generation;
|
||||||
|
`
|
||||||
|
|
||||||
|
args := pgx.StrictNamedArgs{
|
||||||
|
"id": e.ID.String(),
|
||||||
|
"status": e.Status,
|
||||||
|
"result": e.Result,
|
||||||
|
"error_message": e.ErrorMessage,
|
||||||
|
"started_at": e.StartedAt,
|
||||||
|
"lease_expires_at": e.LeaseExpiresAt,
|
||||||
|
"updated_at": e.UpdatedAt,
|
||||||
|
"lease_generation": leaseGeneration,
|
||||||
|
}
|
||||||
|
|
||||||
|
tag, err := tx.Exec(ctx, q, args)
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("cannot commit agent run result: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return tag.RowsAffected(), nil
|
||||||
|
}
|
||||||
|
|
||||||
func (e *AgentRun) LoadNextPendingForUpdateSkipLocked(
|
func (e *AgentRun) LoadNextPendingForUpdateSkipLocked(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
tx pg.Tx,
|
tx pg.Tx,
|
||||||
@@ -531,8 +571,8 @@ SELECT
|
|||||||
result,
|
result,
|
||||||
error_message,
|
error_message,
|
||||||
started_at,
|
started_at,
|
||||||
lease_owner,
|
|
||||||
lease_expires_at,
|
lease_expires_at,
|
||||||
|
lease_generation,
|
||||||
created_at,
|
created_at,
|
||||||
updated_at
|
updated_at
|
||||||
FROM
|
FROM
|
||||||
@@ -566,7 +606,7 @@ FOR UPDATE SKIP LOCKED;
|
|||||||
// ResetStaleAgentRuns resets agent runs whose worker lease has expired.
|
// ResetStaleAgentRuns resets agent runs whose worker lease has expired.
|
||||||
// The worker refreshes lease_expires_at from a separate heartbeat goroutine,
|
// The worker refreshes lease_expires_at from a separate heartbeat goroutine,
|
||||||
// so a long LLM or tool call is not considered stale while the process is alive.
|
// so a long LLM or tool call is not considered stale while the process is alive.
|
||||||
// Stale recovery returns rows to PENDING so the supervisor auto-resumes
|
// Stale recovery returns rows to PENDING so the worker auto-resumes
|
||||||
// from checkpoint when one exists.
|
// from checkpoint when one exists.
|
||||||
func ResetStaleAgentRuns(ctx context.Context, conn pg.Querier) error {
|
func ResetStaleAgentRuns(ctx context.Context, conn pg.Querier) error {
|
||||||
q := `
|
q := `
|
||||||
@@ -574,7 +614,6 @@ UPDATE agent_runs
|
|||||||
SET
|
SET
|
||||||
status = 'PENDING',
|
status = 'PENDING',
|
||||||
started_at = NULL,
|
started_at = NULL,
|
||||||
lease_owner = NULL,
|
|
||||||
lease_expires_at = NULL,
|
lease_expires_at = NULL,
|
||||||
updated_at = now()
|
updated_at = now()
|
||||||
WHERE
|
WHERE
|
||||||
@@ -597,7 +636,7 @@ func HeartbeatAgentRunLease(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
conn pg.Querier,
|
conn pg.Querier,
|
||||||
runID string,
|
runID string,
|
||||||
leaseOwner string,
|
leaseGeneration int64,
|
||||||
expiresAt time.Time,
|
expiresAt time.Time,
|
||||||
) (int64, error) {
|
) (int64, error) {
|
||||||
q := `
|
q := `
|
||||||
@@ -608,13 +647,13 @@ SET
|
|||||||
WHERE
|
WHERE
|
||||||
id = @id
|
id = @id
|
||||||
AND status = 'RUNNING'
|
AND status = 'RUNNING'
|
||||||
AND lease_owner = @lease_owner;
|
AND lease_generation = @lease_generation;
|
||||||
`
|
`
|
||||||
|
|
||||||
args := pgx.StrictNamedArgs{
|
args := pgx.StrictNamedArgs{
|
||||||
"id": runID,
|
"id": runID,
|
||||||
"lease_owner": leaseOwner,
|
|
||||||
"lease_expires_at": expiresAt,
|
"lease_expires_at": expiresAt,
|
||||||
|
"lease_generation": leaseGeneration,
|
||||||
}
|
}
|
||||||
|
|
||||||
tag, err := conn.Exec(ctx, q, args)
|
tag, err := conn.Exec(ctx, q, args)
|
||||||
@@ -699,6 +738,55 @@ WHERE
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *PGCheckpointer) SaveForLease(
|
||||||
|
ctx context.Context,
|
||||||
|
runID string,
|
||||||
|
cp *agent.Checkpoint,
|
||||||
|
leaseGeneration int64,
|
||||||
|
) error {
|
||||||
|
if _, err := gid.ParseGID(runID); err != nil {
|
||||||
|
return fmt.Errorf("cannot parse agent run id: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := s.marshalAgentCheckpoint(cp)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.pg.WithConn(
|
||||||
|
ctx,
|
||||||
|
func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
q := `
|
||||||
|
UPDATE agent_runs
|
||||||
|
SET
|
||||||
|
checkpoint = @checkpoint,
|
||||||
|
updated_at = now()
|
||||||
|
WHERE
|
||||||
|
id = @id
|
||||||
|
AND status = 'RUNNING'
|
||||||
|
AND lease_generation = @lease_generation;
|
||||||
|
`
|
||||||
|
|
||||||
|
args := pgx.StrictNamedArgs{
|
||||||
|
"id": runID,
|
||||||
|
"checkpoint": json.RawMessage(data),
|
||||||
|
"lease_generation": leaseGeneration,
|
||||||
|
}
|
||||||
|
|
||||||
|
tag, err := conn.Exec(ctx, q, args)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("cannot save checkpoint: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if tag.RowsAffected() == 0 {
|
||||||
|
return fmt.Errorf("cannot save checkpoint: lease lost")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
func (s *PGCheckpointer) Load(ctx context.Context, runID string) (*agent.Checkpoint, error) {
|
func (s *PGCheckpointer) Load(ctx context.Context, runID string) (*agent.Checkpoint, error) {
|
||||||
if _, err := gid.ParseGID(runID); err != nil {
|
if _, err := gid.ParseGID(runID); err != nil {
|
||||||
return nil, fmt.Errorf("cannot parse agent run id: %w", err)
|
return nil, fmt.Errorf("cannot parse agent run id: %w", err)
|
||||||
|
|||||||
17
pkg/coredata/migrations/20260607T060000Z.sql
Normal file
17
pkg/coredata/migrations/20260607T060000Z.sql
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
-- Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
--
|
||||||
|
-- Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
-- purpose with or without fee is hereby granted, provided that the above
|
||||||
|
-- copyright notice and this permission notice appear in all copies.
|
||||||
|
--
|
||||||
|
-- THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
-- REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
-- AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
-- INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
-- LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
-- OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
-- PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
ALTER TABLE agent_runs
|
||||||
|
ADD COLUMN lease_generation BIGINT NOT NULL DEFAULT 0,
|
||||||
|
DROP COLUMN lease_owner;
|
||||||
@@ -21,7 +21,6 @@ import (
|
|||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"go.probo.inc/probo/pkg/agent"
|
"go.probo.inc/probo/pkg/agent"
|
||||||
"go.probo.inc/probo/pkg/agentruntest"
|
|
||||||
"go.probo.inc/probo/pkg/coredata"
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
"go.probo.inc/probo/pkg/gid"
|
"go.probo.inc/probo/pkg/gid"
|
||||||
"go.probo.inc/probo/pkg/llm"
|
"go.probo.inc/probo/pkg/llm"
|
||||||
@@ -30,7 +29,7 @@ import (
|
|||||||
func TestPGCheckpointer(t *testing.T) {
|
func TestPGCheckpointer(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
client := agentruntest.PGClient(t)
|
client := pgClient(t)
|
||||||
store := coredata.NewPGCheckpointer(client)
|
store := coredata.NewPGCheckpointer(client)
|
||||||
|
|
||||||
t.Run(
|
t.Run(
|
||||||
@@ -39,7 +38,7 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
run := agentruntest.InsertPendingRun(
|
run := insertPendingRun(
|
||||||
t,
|
t,
|
||||||
client,
|
client,
|
||||||
"test-agent",
|
"test-agent",
|
||||||
@@ -58,7 +57,7 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
run := agentruntest.InsertPendingRun(
|
run := insertPendingRun(
|
||||||
t,
|
t,
|
||||||
client,
|
client,
|
||||||
"test-agent",
|
"test-agent",
|
||||||
@@ -100,7 +99,7 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
run := agentruntest.InsertPendingRun(
|
run := insertPendingRun(
|
||||||
t,
|
t,
|
||||||
client,
|
client,
|
||||||
"test-agent",
|
"test-agent",
|
||||||
@@ -143,7 +142,7 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
run := agentruntest.InsertPendingRun(
|
run := insertPendingRun(
|
||||||
t,
|
t,
|
||||||
client,
|
client,
|
||||||
"test-agent",
|
"test-agent",
|
||||||
@@ -202,7 +201,7 @@ func TestPGCheckpointer(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
run := agentruntest.InsertPendingRun(
|
run := insertPendingRun(
|
||||||
t,
|
t,
|
||||||
client,
|
client,
|
||||||
"test-agent",
|
"test-agent",
|
||||||
|
|||||||
202
pkg/coredata/pg_test_helpers_test.go
Normal file
202
pkg/coredata/pg_test_helpers_test.go
Normal file
@@ -0,0 +1,202 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
// purpose with or without fee is hereby granted, provided that the above
|
||||||
|
// copyright notice and this permission notice appear in all copies.
|
||||||
|
//
|
||||||
|
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
package coredata_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.gearno.de/kit/pg"
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/gid"
|
||||||
|
"go.probo.inc/probo/pkg/llm"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
sharedPGClientCoredata *pg.Client
|
||||||
|
pgOnceCoredata sync.Once
|
||||||
|
pgInitErrCoredata error
|
||||||
|
ensureTableOnceCoredata sync.Once
|
||||||
|
ensureTableErrCoredata error
|
||||||
|
)
|
||||||
|
|
||||||
|
func pgClient(t *testing.T) *pg.Client {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
pgOnceCoredata.Do(func() {
|
||||||
|
addr := os.Getenv("PROBO_TEST_PG_ADDR")
|
||||||
|
if addr == "" {
|
||||||
|
addr = "localhost:5432"
|
||||||
|
}
|
||||||
|
|
||||||
|
user := os.Getenv("PROBO_TEST_PG_USER")
|
||||||
|
if user == "" {
|
||||||
|
user = "probod"
|
||||||
|
}
|
||||||
|
|
||||||
|
password := os.Getenv("PROBO_TEST_PG_PASSWORD")
|
||||||
|
if password == "" {
|
||||||
|
password = "probod"
|
||||||
|
}
|
||||||
|
|
||||||
|
database := os.Getenv("PROBO_TEST_PG_DATABASE")
|
||||||
|
if database == "" {
|
||||||
|
database = "probod_test"
|
||||||
|
}
|
||||||
|
|
||||||
|
sharedPGClientCoredata, pgInitErrCoredata = pg.NewClient(
|
||||||
|
pg.WithAddr(addr),
|
||||||
|
pg.WithUser(user),
|
||||||
|
pg.WithPassword(password),
|
||||||
|
pg.WithDatabase(database),
|
||||||
|
pg.WithPoolSize(5),
|
||||||
|
)
|
||||||
|
if pgInitErrCoredata != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
pgInitErrCoredata = sharedPGClientCoredata.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
_, err := conn.Exec(ctx, "SELECT 1")
|
||||||
|
return err
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
if pgInitErrCoredata != nil {
|
||||||
|
t.Skipf("cannot connect to test database: %v", pgInitErrCoredata)
|
||||||
|
}
|
||||||
|
|
||||||
|
ensureAgentRunsTable(t, sharedPGClientCoredata)
|
||||||
|
|
||||||
|
return sharedPGClientCoredata
|
||||||
|
}
|
||||||
|
|
||||||
|
func ensureAgentRunsTable(t *testing.T, client *pg.Client) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
ensureTableOnceCoredata.Do(func() {
|
||||||
|
ctx := context.Background()
|
||||||
|
ensureTableErrCoredata = client.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
var exists bool
|
||||||
|
if err := conn.QueryRow(
|
||||||
|
ctx,
|
||||||
|
`SELECT EXISTS (SELECT 1 FROM information_schema.tables WHERE table_name = 'agent_runs')`,
|
||||||
|
).Scan(&exists); err != nil {
|
||||||
|
return fmt.Errorf("cannot check agent_runs existence: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !exists {
|
||||||
|
ddl, err := coredata.Migrations.ReadFile("migrations/20260424T173529Z.sql")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("cannot read agent_runs base migration: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := conn.Exec(ctx, string(ddl)); err != nil {
|
||||||
|
return fmt.Errorf("cannot apply agent_runs base migration: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var hasLeaseGeneration bool
|
||||||
|
if err := conn.QueryRow(
|
||||||
|
ctx,
|
||||||
|
`SELECT EXISTS (
|
||||||
|
SELECT 1
|
||||||
|
FROM information_schema.columns
|
||||||
|
WHERE table_name = 'agent_runs'
|
||||||
|
AND column_name = 'lease_generation'
|
||||||
|
)`,
|
||||||
|
).Scan(&hasLeaseGeneration); err != nil {
|
||||||
|
return fmt.Errorf("cannot check lease_generation column: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hasLeaseGeneration {
|
||||||
|
ddl, err := coredata.Migrations.ReadFile("migrations/20260607T060000Z.sql")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("cannot read agent_runs lease generation migration: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := conn.Exec(ctx, string(ddl)); err != nil {
|
||||||
|
return fmt.Errorf("cannot apply agent_runs lease generation migration: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
require.NoError(t, ensureTableErrCoredata, "cannot ensure agent_runs table")
|
||||||
|
}
|
||||||
|
|
||||||
|
func insertPendingRun(
|
||||||
|
t *testing.T,
|
||||||
|
client *pg.Client,
|
||||||
|
agentName string,
|
||||||
|
inputMessages []llm.Message,
|
||||||
|
) coredata.AgentRun {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
tenantID := gid.NewTenantID()
|
||||||
|
orgID := gid.New(tenantID, coredata.OrganizationEntityType)
|
||||||
|
runID := gid.New(tenantID, coredata.AgentRunEntityType)
|
||||||
|
inputJSON, err := json.Marshal(inputMessages)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
run := coredata.AgentRun{
|
||||||
|
ID: runID,
|
||||||
|
OrganizationID: orgID,
|
||||||
|
StartAgentName: agentName,
|
||||||
|
Status: coredata.AgentRunStatusPending,
|
||||||
|
InputMessages: inputJSON,
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
|
}
|
||||||
|
|
||||||
|
err = client.WithTx(
|
||||||
|
context.Background(),
|
||||||
|
func(ctx context.Context, tx pg.Tx) error {
|
||||||
|
if _, err := tx.Exec(
|
||||||
|
ctx,
|
||||||
|
`INSERT INTO organizations (id, tenant_id, name, created_at, updated_at) VALUES ($1, $2, $3, $4, $5)`,
|
||||||
|
orgID.String(), tenantID.String(), "test-org-"+orgID.String(), now, now,
|
||||||
|
); err != nil {
|
||||||
|
return fmt.Errorf("cannot insert placeholder organization: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return run.Insert(ctx, tx, coredata.NewScope(tenantID))
|
||||||
|
},
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
_ = client.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
||||||
|
_, err := conn.Exec(ctx, "DELETE FROM organizations WHERE id = $1", orgID.String())
|
||||||
|
return err
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
return run
|
||||||
|
}
|
||||||
@@ -133,6 +133,10 @@ const (
|
|||||||
ActionThirdPartyRiskAssessmentCreate = "core:thirdParty-risk-assessment:create"
|
ActionThirdPartyRiskAssessmentCreate = "core:thirdParty-risk-assessment:create"
|
||||||
ActionThirdPartyRiskAssessmentList = "core:thirdParty-risk-assessment:list"
|
ActionThirdPartyRiskAssessmentList = "core:thirdParty-risk-assessment:list"
|
||||||
|
|
||||||
|
// AgentRun actions
|
||||||
|
ActionAgentRunGet = "core:agent-run:get"
|
||||||
|
ActionAgentRunList = "core:agent-run:list"
|
||||||
|
|
||||||
// Framework actions
|
// Framework actions
|
||||||
ActionFrameworkGet = "core:framework:get"
|
ActionFrameworkGet = "core:framework:get"
|
||||||
ActionFrameworkList = "core:framework:list"
|
ActionFrameworkList = "core:framework:list"
|
||||||
|
|||||||
@@ -56,6 +56,7 @@ var ViewerPolicy = policy.NewPolicy(
|
|||||||
ActionThirdPartyDataPrivacyAgreementGet,
|
ActionThirdPartyDataPrivacyAgreementGet,
|
||||||
ActionThirdPartyRiskAssessmentList,
|
ActionThirdPartyRiskAssessmentList,
|
||||||
ActionThirdPartyRelationList,
|
ActionThirdPartyRelationList,
|
||||||
|
ActionAgentRunGet, ActionAgentRunList,
|
||||||
ActionFrameworkGet, ActionFrameworkList,
|
ActionFrameworkGet, ActionFrameworkList,
|
||||||
ActionControlGet, ActionControlList,
|
ActionControlGet, ActionControlList,
|
||||||
ActionMeasureGet, ActionMeasureList,
|
ActionMeasureGet, ActionMeasureList,
|
||||||
@@ -144,6 +145,7 @@ var AuditorPolicy = policy.NewPolicy(
|
|||||||
ActionThirdPartyDataPrivacyAgreementGet,
|
ActionThirdPartyDataPrivacyAgreementGet,
|
||||||
ActionThirdPartyRiskAssessmentList,
|
ActionThirdPartyRiskAssessmentList,
|
||||||
ActionThirdPartyRelationList,
|
ActionThirdPartyRelationList,
|
||||||
|
ActionAgentRunGet, ActionAgentRunList,
|
||||||
ActionFrameworkGet, ActionFrameworkList,
|
ActionFrameworkGet, ActionFrameworkList,
|
||||||
ActionControlGet, ActionControlList,
|
ActionControlGet, ActionControlList,
|
||||||
ActionMeasureGet, ActionMeasureList,
|
ActionMeasureGet, ActionMeasureList,
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ import (
|
|||||||
"go.opentelemetry.io/otel/trace"
|
"go.opentelemetry.io/otel/trace"
|
||||||
"go.probo.inc/probo/packages/emails"
|
"go.probo.inc/probo/packages/emails"
|
||||||
"go.probo.inc/probo/pkg/accessreview"
|
"go.probo.inc/probo/pkg/accessreview"
|
||||||
|
"go.probo.inc/probo/pkg/agentrun"
|
||||||
"go.probo.inc/probo/pkg/awsconfig"
|
"go.probo.inc/probo/pkg/awsconfig"
|
||||||
"go.probo.inc/probo/pkg/baseurl"
|
"go.probo.inc/probo/pkg/baseurl"
|
||||||
"go.probo.inc/probo/pkg/certmanager"
|
"go.probo.inc/probo/pkg/certmanager"
|
||||||
@@ -593,6 +594,8 @@ func (impl *Implm) Run(
|
|||||||
l.Named("access-review"),
|
l.Named("access-review"),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
agentRunService := agentrun.NewService(pgClient)
|
||||||
|
|
||||||
thirdPartyService := thirdparty.NewService(pgClient, fileService, thirdPartyVetter)
|
thirdPartyService := thirdparty.NewService(pgClient, fileService, thirdPartyVetter)
|
||||||
riskManagementService := riskmanagement.NewService(pgClient)
|
riskManagementService := riskmanagement.NewService(pgClient)
|
||||||
|
|
||||||
@@ -606,6 +609,7 @@ func (impl *Implm) Run(
|
|||||||
Trust: trustService,
|
Trust: trustService,
|
||||||
ESign: esignService,
|
ESign: esignService,
|
||||||
AccessReview: accessReviewService,
|
AccessReview: accessReviewService,
|
||||||
|
AgentRun: agentRunService,
|
||||||
Mailman: mailmanService,
|
Mailman: mailmanService,
|
||||||
CookieBanner: cookieBannerService,
|
CookieBanner: cookieBannerService,
|
||||||
Geoloc: geolocService,
|
Geoloc: geolocService,
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ import (
|
|||||||
"go.gearno.de/kit/httpserver"
|
"go.gearno.de/kit/httpserver"
|
||||||
"go.gearno.de/kit/log"
|
"go.gearno.de/kit/log"
|
||||||
"go.probo.inc/probo/pkg/accessreview"
|
"go.probo.inc/probo/pkg/accessreview"
|
||||||
|
"go.probo.inc/probo/pkg/agentrun"
|
||||||
"go.probo.inc/probo/pkg/baseurl"
|
"go.probo.inc/probo/pkg/baseurl"
|
||||||
"go.probo.inc/probo/pkg/connector"
|
"go.probo.inc/probo/pkg/connector"
|
||||||
"go.probo.inc/probo/pkg/connector/provider"
|
"go.probo.inc/probo/pkg/connector/provider"
|
||||||
@@ -60,6 +61,7 @@ type (
|
|||||||
Trust *trust.Service
|
Trust *trust.Service
|
||||||
ESign *esign.Service
|
ESign *esign.Service
|
||||||
AccessReview *accessreview.Service
|
AccessReview *accessreview.Service
|
||||||
|
AgentRun *agentrun.Service
|
||||||
Slack *slack.Service
|
Slack *slack.Service
|
||||||
Mailman *mailman.Service
|
Mailman *mailman.Service
|
||||||
CookieBanner *cookiebanner.Service
|
CookieBanner *cookiebanner.Service
|
||||||
@@ -189,6 +191,7 @@ func NewServer(cfg Config) (*Server, error) {
|
|||||||
cfg.IAM,
|
cfg.IAM,
|
||||||
cfg.ESign,
|
cfg.ESign,
|
||||||
cfg.AccessReview,
|
cfg.AccessReview,
|
||||||
|
cfg.AgentRun,
|
||||||
cfg.Mailman,
|
cfg.Mailman,
|
||||||
cfg.CookieBanner,
|
cfg.CookieBanner,
|
||||||
cfg.Cookie,
|
cfg.Cookie,
|
||||||
|
|||||||
82
pkg/server/api/console/v1/agent_run_resolvers.go
Normal file
82
pkg/server/api/console/v1/agent_run_resolvers.go
Normal file
@@ -0,0 +1,82 @@
|
|||||||
|
package console_v1
|
||||||
|
|
||||||
|
// This file will be automatically regenerated based on the schema, any resolver
|
||||||
|
// implementations
|
||||||
|
// will be copied through when generating and any unknown code will be moved to the end.
|
||||||
|
// Code generated by github.com/99designs/gqlgen version v0.17.90
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/vikstrous/dataloadgen"
|
||||||
|
"go.gearno.de/kit/log"
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/probo"
|
||||||
|
"go.probo.inc/probo/pkg/server/api/console/v1/dataloader"
|
||||||
|
"go.probo.inc/probo/pkg/server/api/console/v1/schema"
|
||||||
|
"go.probo.inc/probo/pkg/server/api/console/v1/types"
|
||||||
|
"go.probo.inc/probo/pkg/server/gqlutils"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Organization is the resolver for the organization field.
|
||||||
|
func (r *agentRunResolver) Organization(ctx context.Context, obj *types.AgentRun) (*types.Organization, error) {
|
||||||
|
if _, err := r.authorize(ctx, obj.ID, probo.ActionOrganizationGet); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
loaders := dataloader.FromContext(ctx)
|
||||||
|
|
||||||
|
organization, err := loaders.Organization.Load(ctx, obj.Organization.ID)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, coredata.ErrResourceNotFound) || errors.Is(err, dataloadgen.ErrNotFound) {
|
||||||
|
return nil, gqlutils.NotFound(ctx, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.logger.ErrorCtx(ctx, "cannot load organization", log.Error(err))
|
||||||
|
|
||||||
|
return nil, gqlutils.Internal(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
return types.NewOrganization(organization), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Permission is the resolver for the permission field.
|
||||||
|
func (r *agentRunResolver) Permission(ctx context.Context, obj *types.AgentRun, action string) (bool, error) {
|
||||||
|
return r.Resolver.Permission(ctx, obj, action)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TotalCount is the resolver for the totalCount field.
|
||||||
|
func (r *agentRunConnectionResolver) TotalCount(ctx context.Context, obj *types.AgentRunConnection) (int, error) {
|
||||||
|
scope, err := r.authorize(ctx, obj.ParentID, probo.ActionAgentRunList)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch obj.Resolver.(type) {
|
||||||
|
case *organizationResolver:
|
||||||
|
count, err := r.agentRun.CountForOrganizationID(ctx, scope, obj.ParentID)
|
||||||
|
if err != nil {
|
||||||
|
r.logger.ErrorCtx(ctx, "cannot count agent runs", log.Error(err))
|
||||||
|
return 0, gqlutils.Internal(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
r.logger.ErrorCtx(ctx, "unsupported resolver for agent run connection", log.String("resolver", fmt.Sprintf("%T", obj.Resolver)))
|
||||||
|
|
||||||
|
return 0, gqlutils.Internal(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AgentRun returns schema.AgentRunResolver implementation.
|
||||||
|
func (r *Resolver) AgentRun() schema.AgentRunResolver { return &agentRunResolver{r} }
|
||||||
|
|
||||||
|
// AgentRunConnection returns schema.AgentRunConnectionResolver implementation.
|
||||||
|
func (r *Resolver) AgentRunConnection() schema.AgentRunConnectionResolver {
|
||||||
|
return &agentRunConnectionResolver{r}
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentRunResolver struct{ *Resolver }
|
||||||
|
type agentRunConnectionResolver struct{ *Resolver }
|
||||||
@@ -369,6 +369,16 @@ func (r *queryResolver) Node(ctx context.Context, id gid.GID) (types.Node, error
|
|||||||
|
|
||||||
return types.NewWebhookSubscription(wc), nil
|
return types.NewWebhookSubscription(wc), nil
|
||||||
}
|
}
|
||||||
|
case coredata.AgentRunEntityType:
|
||||||
|
action = probo.ActionAgentRunGet
|
||||||
|
loadNode = func(ctx context.Context, scope *coredata.Scope, id gid.GID) (types.Node, error) {
|
||||||
|
run, err := r.agentRun.Get(ctx, scope, id)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return types.NewAgentRun(run), nil
|
||||||
|
}
|
||||||
case coredata.AccessReviewCampaignEntityType:
|
case coredata.AccessReviewCampaignEntityType:
|
||||||
action = probo.ActionAccessReviewCampaignGet
|
action = probo.ActionAccessReviewCampaignGet
|
||||||
loadNode = func(ctx context.Context, scope *coredata.Scope, id gid.GID) (types.Node, error) {
|
loadNode = func(ctx context.Context, scope *coredata.Scope, id gid.GID) (types.Node, error) {
|
||||||
|
|||||||
60
pkg/server/api/console/v1/graphql/agent_run.graphql
Normal file
60
pkg/server/api/console/v1/graphql/agent_run.graphql
Normal file
@@ -0,0 +1,60 @@
|
|||||||
|
enum AgentRunStatus
|
||||||
|
@goModel(model: "go.probo.inc/probo/pkg/coredata.AgentRunStatus") {
|
||||||
|
PENDING
|
||||||
|
@goEnum(value: "go.probo.inc/probo/pkg/coredata.AgentRunStatusPending")
|
||||||
|
RUNNING
|
||||||
|
@goEnum(value: "go.probo.inc/probo/pkg/coredata.AgentRunStatusRunning")
|
||||||
|
SUSPENDED
|
||||||
|
@goEnum(value: "go.probo.inc/probo/pkg/coredata.AgentRunStatusSuspended")
|
||||||
|
AWAITING_APPROVAL
|
||||||
|
@goEnum(
|
||||||
|
value: "go.probo.inc/probo/pkg/coredata.AgentRunStatusAwaitingApproval"
|
||||||
|
)
|
||||||
|
COMPLETED
|
||||||
|
@goEnum(value: "go.probo.inc/probo/pkg/coredata.AgentRunStatusCompleted")
|
||||||
|
FAILED
|
||||||
|
@goEnum(value: "go.probo.inc/probo/pkg/coredata.AgentRunStatusFailed")
|
||||||
|
}
|
||||||
|
|
||||||
|
enum AgentRunOrderField
|
||||||
|
@goModel(model: "go.probo.inc/probo/pkg/coredata.AgentRunOrderField") {
|
||||||
|
CREATED_AT
|
||||||
|
@goEnum(
|
||||||
|
value: "go.probo.inc/probo/pkg/coredata.AgentRunOrderFieldCreatedAt"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
input AgentRunOrder
|
||||||
|
@goModel(
|
||||||
|
model: "go.probo.inc/probo/pkg/server/api/console/v1/types.AgentRunOrderBy"
|
||||||
|
) {
|
||||||
|
direction: OrderDirection!
|
||||||
|
field: AgentRunOrderField!
|
||||||
|
}
|
||||||
|
|
||||||
|
type AgentRun implements Node {
|
||||||
|
id: ID!
|
||||||
|
organization: Organization! @goField(forceResolver: true)
|
||||||
|
agentName: String!
|
||||||
|
status: AgentRunStatus!
|
||||||
|
errorMessage: String
|
||||||
|
startedAt: Datetime
|
||||||
|
createdAt: Datetime!
|
||||||
|
updatedAt: Datetime!
|
||||||
|
|
||||||
|
permission(action: String!): Boolean! @goField(forceResolver: true)
|
||||||
|
}
|
||||||
|
|
||||||
|
type AgentRunConnection
|
||||||
|
@goModel(
|
||||||
|
model: "go.probo.inc/probo/pkg/server/api/console/v1/types.AgentRunConnection"
|
||||||
|
) {
|
||||||
|
totalCount: Int! @goField(forceResolver: true)
|
||||||
|
edges: [AgentRunEdge!]!
|
||||||
|
pageInfo: PageInfo!
|
||||||
|
}
|
||||||
|
|
||||||
|
type AgentRunEdge {
|
||||||
|
cursor: CursorKey!
|
||||||
|
node: AgentRun!
|
||||||
|
}
|
||||||
@@ -311,6 +311,14 @@ type Organization implements Node {
|
|||||||
orderBy: TaskOrder
|
orderBy: TaskOrder
|
||||||
): TaskConnection! @goField(forceResolver: true)
|
): TaskConnection! @goField(forceResolver: true)
|
||||||
|
|
||||||
|
agentRuns(
|
||||||
|
first: Int
|
||||||
|
after: CursorKey
|
||||||
|
last: Int
|
||||||
|
before: CursorKey
|
||||||
|
orderBy: AgentRunOrder
|
||||||
|
): AgentRunConnection! @goField(forceResolver: true)
|
||||||
|
|
||||||
trustCenter: TrustCenter @goField(forceResolver: true)
|
trustCenter: TrustCenter @goField(forceResolver: true)
|
||||||
customDomain: CustomDomain @goField(forceResolver: true)
|
customDomain: CustomDomain @goField(forceResolver: true)
|
||||||
trustCenterFiles(
|
trustCenterFiles(
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ import (
|
|||||||
|
|
||||||
"go.gearno.de/kit/log"
|
"go.gearno.de/kit/log"
|
||||||
"go.probo.inc/probo/pkg/accessreview"
|
"go.probo.inc/probo/pkg/accessreview"
|
||||||
|
"go.probo.inc/probo/pkg/agentrun"
|
||||||
"go.probo.inc/probo/pkg/connector"
|
"go.probo.inc/probo/pkg/connector"
|
||||||
"go.probo.inc/probo/pkg/connector/provider"
|
"go.probo.inc/probo/pkg/connector/provider"
|
||||||
"go.probo.inc/probo/pkg/cookiebanner"
|
"go.probo.inc/probo/pkg/cookiebanner"
|
||||||
@@ -39,6 +40,7 @@ func NewGraphQLHandler(
|
|||||||
proboSvc *probo.Service,
|
proboSvc *probo.Service,
|
||||||
esignSvc *esign.Service,
|
esignSvc *esign.Service,
|
||||||
accessReviewSvc *accessreview.Service,
|
accessReviewSvc *accessreview.Service,
|
||||||
|
agentRunSvc *agentrun.Service,
|
||||||
mailmanSvc *mailman.Service,
|
mailmanSvc *mailman.Service,
|
||||||
cookieBannerSvc *cookiebanner.Service,
|
cookieBannerSvc *cookiebanner.Service,
|
||||||
connectorRegistry *connector.ConnectorRegistry,
|
connectorRegistry *connector.ConnectorRegistry,
|
||||||
@@ -56,6 +58,7 @@ func NewGraphQLHandler(
|
|||||||
iam: iamSvc,
|
iam: iamSvc,
|
||||||
esign: esignSvc,
|
esign: esignSvc,
|
||||||
accessReview: accessReviewSvc,
|
accessReview: accessReviewSvc,
|
||||||
|
agentRun: agentRunSvc,
|
||||||
mailman: mailmanSvc,
|
mailman: mailmanSvc,
|
||||||
cookieBanner: cookieBannerSvc,
|
cookieBanner: cookieBannerSvc,
|
||||||
connectorRegistry: connectorRegistry,
|
connectorRegistry: connectorRegistry,
|
||||||
|
|||||||
@@ -1173,6 +1173,36 @@ func (r *organizationResolver) Tasks(ctx context.Context, obj *types.Organizatio
|
|||||||
return types.NewTaskConnection(page, r, obj.ID), nil
|
return types.NewTaskConnection(page, r, obj.ID), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AgentRuns is the resolver for the agentRuns field.
|
||||||
|
func (r *organizationResolver) AgentRuns(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.AgentRunOrderBy) (*types.AgentRunConnection, error) {
|
||||||
|
scope, err := r.authorize(ctx, obj.ID, probo.ActionAgentRunList)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
pageOrderBy := page.OrderBy[coredata.AgentRunOrderField]{
|
||||||
|
Field: coredata.AgentRunOrderFieldCreatedAt,
|
||||||
|
Direction: page.OrderDirectionDesc,
|
||||||
|
}
|
||||||
|
|
||||||
|
if orderBy != nil {
|
||||||
|
pageOrderBy = page.OrderBy[coredata.AgentRunOrderField]{
|
||||||
|
Field: orderBy.Field,
|
||||||
|
Direction: orderBy.Direction,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
|
||||||
|
|
||||||
|
page, err := r.agentRun.ListForOrganizationID(ctx, scope, obj.ID, cursor)
|
||||||
|
if err != nil {
|
||||||
|
r.logger.ErrorCtx(ctx, "cannot list organization agent runs", log.Error(err))
|
||||||
|
return nil, gqlutils.Internal(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
return types.NewAgentRunConnection(page, r, obj.ID), nil
|
||||||
|
}
|
||||||
|
|
||||||
// TrustCenter is the resolver for the trustCenter field.
|
// TrustCenter is the resolver for the trustCenter field.
|
||||||
func (r *organizationResolver) TrustCenter(ctx context.Context, obj *types.Organization) (*types.TrustCenter, error) {
|
func (r *organizationResolver) TrustCenter(ctx context.Context, obj *types.Organization) (*types.TrustCenter, error) {
|
||||||
scope, err := r.authorize(ctx, obj.ID, probo.ActionTrustCenterGet)
|
scope, err := r.authorize(ctx, obj.ID, probo.ActionTrustCenterGet)
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ import (
|
|||||||
"go.gearno.de/kit/httpserver"
|
"go.gearno.de/kit/httpserver"
|
||||||
"go.gearno.de/kit/log"
|
"go.gearno.de/kit/log"
|
||||||
"go.probo.inc/probo/pkg/accessreview"
|
"go.probo.inc/probo/pkg/accessreview"
|
||||||
|
"go.probo.inc/probo/pkg/agentrun"
|
||||||
"go.probo.inc/probo/pkg/baseurl"
|
"go.probo.inc/probo/pkg/baseurl"
|
||||||
"go.probo.inc/probo/pkg/connector"
|
"go.probo.inc/probo/pkg/connector"
|
||||||
"go.probo.inc/probo/pkg/connector/provider"
|
"go.probo.inc/probo/pkg/connector/provider"
|
||||||
@@ -55,6 +56,7 @@ type (
|
|||||||
iam *iam.Service
|
iam *iam.Service
|
||||||
esign *esign.Service
|
esign *esign.Service
|
||||||
accessReview *accessreview.Service
|
accessReview *accessreview.Service
|
||||||
|
agentRun *agentrun.Service
|
||||||
mailman *mailman.Service
|
mailman *mailman.Service
|
||||||
cookieBanner *cookiebanner.Service
|
cookieBanner *cookiebanner.Service
|
||||||
connectorRegistry *connector.ConnectorRegistry
|
connectorRegistry *connector.ConnectorRegistry
|
||||||
@@ -72,6 +74,7 @@ func NewMux(
|
|||||||
iamSvc *iam.Service,
|
iamSvc *iam.Service,
|
||||||
esignSvc *esign.Service,
|
esignSvc *esign.Service,
|
||||||
accessReviewSvc *accessreview.Service,
|
accessReviewSvc *accessreview.Service,
|
||||||
|
agentRunSvc *agentrun.Service,
|
||||||
mailmanSvc *mailman.Service,
|
mailmanSvc *mailman.Service,
|
||||||
cookieBannerSvc *cookiebanner.Service,
|
cookieBannerSvc *cookiebanner.Service,
|
||||||
cookieConfig securecookie.Config,
|
cookieConfig securecookie.Config,
|
||||||
@@ -92,6 +95,7 @@ func NewMux(
|
|||||||
proboSvc,
|
proboSvc,
|
||||||
esignSvc,
|
esignSvc,
|
||||||
accessReviewSvc,
|
accessReviewSvc,
|
||||||
|
agentRunSvc,
|
||||||
mailmanSvc,
|
mailmanSvc,
|
||||||
cookieBannerSvc,
|
cookieBannerSvc,
|
||||||
connectorRegistry,
|
connectorRegistry,
|
||||||
|
|||||||
76
pkg/server/api/console/v1/types/agent_run.go
Normal file
76
pkg/server/api/console/v1/types/agent_run.go
Normal file
@@ -0,0 +1,76 @@
|
|||||||
|
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||||
|
//
|
||||||
|
// Permission to use, copy, modify, and/or distribute this software for any
|
||||||
|
// purpose with or without fee is hereby granted, provided that the above
|
||||||
|
// copyright notice and this permission notice appear in all copies.
|
||||||
|
//
|
||||||
|
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||||
|
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||||
|
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||||
|
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||||
|
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||||
|
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||||
|
// PERFORMANCE OF THIS SOFTWARE.
|
||||||
|
|
||||||
|
package types
|
||||||
|
|
||||||
|
import (
|
||||||
|
"go.probo.inc/probo/pkg/coredata"
|
||||||
|
"go.probo.inc/probo/pkg/gid"
|
||||||
|
"go.probo.inc/probo/pkg/page"
|
||||||
|
)
|
||||||
|
|
||||||
|
type (
|
||||||
|
AgentRunOrderBy OrderBy[coredata.AgentRunOrderField]
|
||||||
|
|
||||||
|
AgentRunConnection struct {
|
||||||
|
TotalCount int
|
||||||
|
Edges []*AgentRunEdge
|
||||||
|
PageInfo PageInfo
|
||||||
|
|
||||||
|
Resolver any
|
||||||
|
ParentID gid.GID
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
func NewAgentRunConnection(
|
||||||
|
p *page.Page[*coredata.AgentRun, coredata.AgentRunOrderField],
|
||||||
|
parentType any,
|
||||||
|
parentID gid.GID,
|
||||||
|
) *AgentRunConnection {
|
||||||
|
var edges = make([]*AgentRunEdge, len(p.Data))
|
||||||
|
|
||||||
|
for i := range edges {
|
||||||
|
edges[i] = NewAgentRunEdge(p.Data[i], p.Cursor.OrderBy.Field)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &AgentRunConnection{
|
||||||
|
Edges: edges,
|
||||||
|
PageInfo: *NewPageInfo(p),
|
||||||
|
|
||||||
|
Resolver: parentType,
|
||||||
|
ParentID: parentID,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewAgentRunEdge(run *coredata.AgentRun, orderBy coredata.AgentRunOrderField) *AgentRunEdge {
|
||||||
|
return &AgentRunEdge{
|
||||||
|
Cursor: run.CursorKey(orderBy),
|
||||||
|
Node: NewAgentRun(run),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewAgentRun(run *coredata.AgentRun) *AgentRun {
|
||||||
|
return &AgentRun{
|
||||||
|
ID: run.ID,
|
||||||
|
Organization: &Organization{
|
||||||
|
ID: run.OrganizationID,
|
||||||
|
},
|
||||||
|
AgentName: run.StartAgentName,
|
||||||
|
Status: run.Status,
|
||||||
|
ErrorMessage: run.ErrorMessage,
|
||||||
|
StartedAt: run.StartedAt,
|
||||||
|
CreatedAt: run.CreatedAt,
|
||||||
|
UpdatedAt: run.UpdatedAt,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -25,6 +25,7 @@ import (
|
|||||||
"go.gearno.de/kit/log"
|
"go.gearno.de/kit/log"
|
||||||
"go.gearno.de/x/ref"
|
"go.gearno.de/x/ref"
|
||||||
"go.probo.inc/probo/pkg/accessreview"
|
"go.probo.inc/probo/pkg/accessreview"
|
||||||
|
"go.probo.inc/probo/pkg/agentrun"
|
||||||
"go.probo.inc/probo/pkg/baseurl"
|
"go.probo.inc/probo/pkg/baseurl"
|
||||||
"go.probo.inc/probo/pkg/connector"
|
"go.probo.inc/probo/pkg/connector"
|
||||||
"go.probo.inc/probo/pkg/connector/provider"
|
"go.probo.inc/probo/pkg/connector/provider"
|
||||||
@@ -59,6 +60,7 @@ type Config struct {
|
|||||||
Trust *trust.Service
|
Trust *trust.Service
|
||||||
ESign *esign.Service
|
ESign *esign.Service
|
||||||
AccessReview *accessreview.Service
|
AccessReview *accessreview.Service
|
||||||
|
AgentRun *agentrun.Service
|
||||||
Slack *slack.Service
|
Slack *slack.Service
|
||||||
Mailman *mailman.Service
|
Mailman *mailman.Service
|
||||||
CookieBanner *cookiebanner.Service
|
CookieBanner *cookiebanner.Service
|
||||||
@@ -97,6 +99,7 @@ func NewServer(cfg Config) (*Server, error) {
|
|||||||
Trust: cfg.Trust,
|
Trust: cfg.Trust,
|
||||||
ESign: cfg.ESign,
|
ESign: cfg.ESign,
|
||||||
AccessReview: cfg.AccessReview,
|
AccessReview: cfg.AccessReview,
|
||||||
|
AgentRun: cfg.AgentRun,
|
||||||
Slack: cfg.Slack,
|
Slack: cfg.Slack,
|
||||||
Mailman: cfg.Mailman,
|
Mailman: cfg.Mailman,
|
||||||
CookieBanner: cfg.CookieBanner,
|
CookieBanner: cfg.CookieBanner,
|
||||||
|
|||||||
Reference in New Issue
Block a user