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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -24,12 +24,13 @@ require (
github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674
github.com/lib/pq v1.10.9
github.com/moby/moby/client v0.2.2
github.com/modelcontextprotocol/go-sdk v1.6.1
github.com/nicholas-fedor/shoutrrr v0.16.1
github.com/robfig/cron/v3 v3.0.1
github.com/spf13/cobra v1.10.2
github.com/testcontainers/testcontainers-go v0.41.0
github.com/testcontainers/testcontainers-go/modules/compose v0.41.0
github.com/whilesmartgo/agents v0.1.0
github.com/whilesmartgo/agents v0.4.0
github.com/whilesmartgo/mcp v0.2.0
go.opentelemetry.io/otel v1.41.0
go.opentelemetry.io/otel/exporters/otlp/otlpmetric/otlpmetricgrpc v1.41.0
Expand Down Expand Up @@ -148,7 +149,6 @@ require (
github.com/moby/sys/user v0.4.0 // indirect
github.com/moby/sys/userns v0.1.0 // indirect
github.com/moby/term v0.5.2 // indirect
github.com/modelcontextprotocol/go-sdk v1.6.1 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee // indirect
github.com/morikuni/aec v1.1.0 // indirect
Expand Down
4 changes: 2 additions & 2 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -518,8 +518,8 @@ github.com/ugorji/go/codec v1.3.0 h1:Qd2W2sQawAfG8XSvzwhBeoGq71zXOC/Q1E9y/wUcsUA
github.com/ugorji/go/codec v1.3.0/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4=
github.com/vbatts/tar-split v0.12.2 h1:w/Y6tjxpeiFMR47yzZPlPj/FcPLpXbTUi/9H7d3CPa4=
github.com/vbatts/tar-split v0.12.2/go.mod h1:eF6B6i6ftWQcDqEn3/iGFRFRo8cBIMSJVOpnNdfTMFA=
github.com/whilesmartgo/agents v0.1.0 h1:AX7Q5e6BY2Fb/1J7prtZgXSx9bsD8GRSDsmfTQ8kgxs=
github.com/whilesmartgo/agents v0.1.0/go.mod h1:MTuXbfen/M5GGphQk3kggnO52i8thPBf3RfywfuxOOM=
github.com/whilesmartgo/agents v0.4.0 h1:DEzeO28Kh31aixHsz55Sq7guYooKEOAUmN2J2VbevzY=
github.com/whilesmartgo/agents v0.4.0/go.mod h1:MTuXbfen/M5GGphQk3kggnO52i8thPBf3RfywfuxOOM=
github.com/whilesmartgo/mcp v0.2.0 h1:oXEI2dagW9rFuRvK4U87iLPd7pxZ08RdAi1HSy8hyrY=
github.com/whilesmartgo/mcp v0.2.0/go.mod h1:GD0Gg3H8rJOSfCQ0o68kXjOg5Kv3QqkAHRwtV4L+xKk=
github.com/xhit/go-str2duration/v2 v2.1.0 h1:lxklc02Drh6ynqX+DdPyp5pCKLUQpRT8bp8Ydu2Bstc=
Expand Down
6 changes: 4 additions & 2 deletions internal/ai/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ func (p *openAICompatible) Name() string {

func (p *openAICompatible) Complete(ctx context.Context, req Request) (*Response, error) {
resp, err := p.engine.Complete(ctx, agents.Request{
Messages: toEngineMessages(req.Messages),
Messages: MessagesToAgents(req.Messages),
Tools: toEngineTools(req.Tools),
MaxTokens: req.MaxTokens,
Temperature: req.Temperature,
Expand All @@ -55,7 +55,9 @@ func (p *openAICompatible) Complete(ctx context.Context, req Request) (*Response
}, nil
}

func toEngineMessages(messages []Message) []agents.Message {
// MessagesToAgents converts the stored transcript to the library's message type.
// Display and Hidden are UI-only and dropped here: the model sees Content.
func MessagesToAgents(messages []Message) []agents.Message {
out := make([]agents.Message, 0, len(messages))
for _, m := range messages {
out = append(out, agents.Message{
Expand Down
76 changes: 76 additions & 0 deletions internal/ai/runner.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
package ai

import (
"context"

"github.com/whilesmartgo/agents"
)

// MessageFromAgents converts a message the runner appended back to the stored
// type. Runner-created messages (assistant and tool turns) carry no UI fields.
func MessageFromAgents(m agents.Message) Message {
return Message{
Role: m.Role,
Content: m.Content,
ToolCalls: fromEngineToolCalls(m.ToolCalls),
ToolCallID: m.ToolCallID,
Name: m.Name,
}
}

// ToolCallsFromAgents converts the runner's tool calls to the stored type.
func ToolCallsFromAgents(calls []agents.ToolCall) []ToolCall {
return fromEngineToolCalls(calls)
}

// CapturingEngine adapts a Provider to the library's Engine. It records the
// model each response reports, which the runner does not otherwise surface, so
// a session can still note which model answered.
type CapturingEngine struct {
provider Provider
lastModel string
}

func NewCapturingEngine(p Provider) *CapturingEngine {
return &CapturingEngine{provider: p}
}

func (e *CapturingEngine) Complete(ctx context.Context, req agents.Request) (*agents.Response, error) {
resp, err := e.provider.Complete(ctx, Request{
Messages: messagesFromAgents(req.Messages),
Tools: toolsFromAgents(req.Tools),
MaxTokens: req.MaxTokens,
Temperature: req.Temperature,
})
if err != nil {
return nil, err
}
if resp.Model != "" {
e.lastModel = resp.Model
}
return &agents.Response{
Content: resp.Content,
ToolCalls: toEngineToolCalls(resp.ToolCalls),
Model: resp.Model,
Usage: agents.Usage{PromptTokens: resp.Usage.PromptTokens, CompletionTokens: resp.Usage.CompletionTokens},
}, nil
}

// LastModel is the model of the most recent response, or "" if none.
func (e *CapturingEngine) LastModel() string { return e.lastModel }

func messagesFromAgents(messages []agents.Message) []Message {
out := make([]Message, 0, len(messages))
for _, m := range messages {
out = append(out, MessageFromAgents(m))
}
return out
}

func toolsFromAgents(schemas []agents.ToolSchema) []Tool {
out := make([]Tool, 0, len(schemas))
for _, s := range schemas {
out = append(out, Tool{Name: s.Name, Description: s.Description, Parameters: s.Parameters})
}
return out
}
137 changes: 98 additions & 39 deletions internal/api/ai_session_handlers.go
Original file line number Diff line number Diff line change
@@ -1,13 +1,17 @@
package api

import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"

"github.com/flatrun/agent/internal/ai"
"github.com/flatrun/agent/internal/auth"
"github.com/gin-gonic/gin"
"github.com/whilesmartgo/agents"
)

// composeUserMessage merges a short message with optional bulky
Expand Down Expand Up @@ -52,42 +56,96 @@ func canUseSession(c *gin.Context, sess *ai.Session) bool {
return sessionActorFrom(c).ID == sess.CreatedBy.ID
}

// advanceSession runs the tool loop: it calls the model, executes any
// requested tools (auto-run) or pauses for approval, and repeats until
// the model returns a final answer or the step budget is exhausted.
const aiStepLimitMessage = "I stopped after investigating several steps without reaching a confident answer. Ask a more specific question or check the details directly."

// advanceSession drives the assistant's tool loop through the shared agents
// runner: it calls the model, runs any requested tools (auto-run) or pauses for
// approval, and repeats until a final answer or the step budget is spent.
func (s *Server) advanceSession(c *gin.Context, sess *ai.Session) error {
tools := s.aiToolSpecs()
for step := 0; step < sess.MaxToolSteps(); step++ {
resp, err := s.aiProvider.Complete(c.Request.Context(), ai.Request{Messages: sess.Messages, Tools: tools})
if err != nil {
return err
}
if sess.Model == "" {
sess.Model = resp.Model
}
engine := ai.NewCapturingEngine(s.aiProvider)
runner := s.aiRunner(c, sess, engine)
conv := &agents.Conversation{Messages: ai.MessagesToAgents(sess.Messages)}
from := len(conv.Messages)
_, err := runner.Advance(c.Request.Context(), conv)
return s.absorbAdvance(sess, conv, from, engine, err)
}

if len(resp.ToolCalls) == 0 {
analysis, suggestions := ai.ParseSuggestions(resp.Content)
sess.AddAssistantMessage(analysis, nil)
sess.Suggested = s.scopeSuggestions(sess, suggestions)
sess.Status = ai.SessionStatusReady
return nil
}
// aiRunner assembles the runner for one session turn. A session that does not
// auto-run pauses before any tools run, surfacing them for per-call approval.
func (s *Server) aiRunner(c *gin.Context, sess *ai.Session, engine agents.Engine) agents.Runner {
runner := agents.Runner{
Engine: engine,
Harness: agents.Harness{
Model: sess.Model,
MaxSteps: sess.MaxToolSteps(),
Tools: s.sessionToolRegistry(c, sess.Deployment),
StepLimitMessage: aiStepLimitMessage,
},
}
if !sess.AutoRun {
runner.Approve = func(context.Context, []agents.ToolCall) (bool, error) { return false, nil }
}
return runner
}

sess.AddAssistantMessage(resp.Content, resp.ToolCalls)
// sessionToolRegistry exposes the assistant's tools to the runner, each bound to
// this request and the session's deployment so per-tool permission and
// protected-mode checks run exactly as they do for a direct tool call.
func (s *Server) sessionToolRegistry(c *gin.Context, deployment string) *agents.Registry {
specs := s.aiToolSpecs()
tools := make([]agents.Tool, 0, len(specs))
for _, spec := range specs {
spec := spec
tools = append(tools, agents.Tool{
Name: spec.Name,
Description: spec.Description,
Parameters: spec.Parameters,
Handler: func(_ context.Context, raw json.RawMessage) (string, error) {
return s.runAITool(c, deployment, ai.ToolCall{Name: spec.Name, Arguments: string(raw)}), nil
},
})
}
return agents.NewRegistry(tools...)
}

if !sess.AutoRun {
sess.Pending = resp.ToolCalls
sess.Status = ai.SessionStatusAwaitingApproval
return nil
}
// absorbAdvance folds the messages the runner appended back into the stored
// session, records the model, and sets the resulting status. A paused turn
// records its pending calls; a real engine error is returned so the caller
// leaves the session unsaved.
func (s *Server) absorbAdvance(sess *ai.Session, conv *agents.Conversation, from int, engine *ai.CapturingEngine, err error) error {
if err != nil && !errors.Is(err, agents.ErrAwaitingApproval) {
return err
}

for _, call := range resp.ToolCalls {
sess.AddToolResult(call, s.runAITool(c, sess.Deployment, call))
added := conv.Messages[from:]
final := err == nil
for i, m := range added {
switch m.Role {
case agents.RoleAssistant:
if i == len(added)-1 && final {
// The last assistant turn of a completed run is the answer; its
// suggestion block is parsed out and offered as one-click actions.
analysis, suggestions := ai.ParseSuggestions(m.Content)
sess.AddAssistantMessage(analysis, ai.ToolCallsFromAgents(m.ToolCalls))
sess.Suggested = s.scopeSuggestions(sess, suggestions)
} else {
sess.AddAssistantMessage(m.Content, ai.ToolCallsFromAgents(m.ToolCalls))
}
case agents.RoleTool:
sess.AddToolResult(ai.ToolCall{ID: m.ToolCallID, Name: m.Name}, m.Content)
}
}
if sess.Model == "" {
sess.Model = engine.LastModel()
}

sess.AddAssistantMessage("I stopped after investigating several steps without reaching a confident answer. Ask a more specific question or check the details directly.", nil)
if errors.Is(err, agents.ErrAwaitingApproval) {
last := conv.Messages[len(conv.Messages)-1]
sess.Pending = ai.ToolCallsFromAgents(last.ToolCalls)
sess.Status = ai.SessionStatusAwaitingApproval
return nil
}
sess.Pending = nil
sess.Status = ai.SessionStatusReady
return nil
}
Expand Down Expand Up @@ -282,17 +340,18 @@ func (s *Server) approveAISessionTools(c *gin.Context) {
return
}

for _, call := range sess.Pending {
if req.Approved[call.ID] {
sess.AddToolResult(call, s.runAITool(c, sess.Deployment, call))
} else {
sess.AddToolResult(call, "The operator declined to run this command.")
}
}
sess.Pending = nil
sess.Status = ai.SessionStatusReady

if err := s.advanceSession(c, sess); err != nil {
// A missing or false decision means declined, so a nil map must not be read
// as "approve all"; an empty non-nil map declines every pending call.
decisions := req.Approved
if decisions == nil {
decisions = map[string]bool{}
}
engine := ai.NewCapturingEngine(s.aiProvider)
runner := s.aiRunner(c, sess, engine)
conv := &agents.Conversation{Messages: ai.MessagesToAgents(sess.Messages)}
from := len(conv.Messages)
_, err := runner.Resume(c.Request.Context(), conv, decisions)
if err := s.absorbAdvance(sess, conv, from, engine, err); err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
Expand Down
49 changes: 49 additions & 0 deletions internal/api/ai_session_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -256,6 +256,55 @@ func TestAISessionDeclineTool(t *testing.T) {
}
}

func TestAISessionPerCallApproval(t *testing.T) {
s, tmpDir, ts := setupPlanTestServer(t)
createTestDeployment(t, tmpDir, "myapp", &models.ServiceMetadata{Name: "myapp"})

s.aiProvider = &scriptedProvider{responses: []*ai.Response{
{ToolCalls: []ai.ToolCall{
{ID: "c1", Name: "list_networks", Arguments: "{}"},
{ID: "c2", Name: "list_deployments", Arguments: "{}"},
}, Model: "scripted"},
{Content: "## Summary\nDone.", Model: "scripted"},
}}

_, parsed := doJSON(t, http.MethodPost, ts.URL+"/api/ai/sessions", map[string]interface{}{
"scope": "system", "auto_run": false, "message": "inspect networks and deployments",
})
id := parsed["id"].(string)
if parsed["status"] != "awaiting_approval" {
t.Fatalf("status = %v, want awaiting_approval", parsed["status"])
}

// Approve one call, decline the other in the same decision.
resp, parsed := doJSON(t, http.MethodPost, ts.URL+"/api/ai/sessions/"+id+"/approve",
map[string]interface{}{"approved": map[string]bool{"c1": true, "c2": false}})
if resp.StatusCode != http.StatusOK {
t.Fatalf("approve status = %d, body %v", resp.StatusCode, parsed)
}
if parsed["status"] != "ready" {
t.Errorf("status = %v, want ready", parsed["status"])
}

results := map[string]string{}
for _, m := range parsed["messages"].([]interface{}) {
steps, ok := m.(map[string]interface{})["tool_steps"].([]interface{})
if !ok {
continue
}
for _, st := range steps {
step := st.(map[string]interface{})
results[step["name"].(string)], _ = step["result"].(string)
}
}
if results["list_networks"] == "" || strings.Contains(results["list_networks"], "declined") {
t.Errorf("approved tool should have run, got %q", results["list_networks"])
}
if !strings.Contains(results["list_deployments"], "declined") {
t.Errorf("declined tool should record the refusal, got %q", results["list_deployments"])
}
}

func TestAISessionDisabledReturns503(t *testing.T) {
_, _, ts := setupPlanTestServer(t)
resp, parsed := doJSON(t, http.MethodPost, ts.URL+"/api/ai/sessions",
Expand Down
Loading