Skip to content
Open
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
2 changes: 1 addition & 1 deletion cmd/ask/agent_provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -172,6 +172,7 @@ func setupAgentSessionTools(s *agentSession, cfg askConfig) {
agentEndTurnTool(env),
agentSearchToolsTool(s.deferredTools),
agentInvokeToolTool(s.deferredTools, s.isCoreToolName, env),
agentLoadArtifactsTool(),
}
if !s.args.InWorkflow {
s.coreTools = append(s.coreTools, agentFinalizedPlanTool(env))
Expand All @@ -183,7 +184,6 @@ func setupAgentSessionTools(s *agentSession, cfg askConfig) {
s.coreTools = append(s.coreTools, agentWorkflowTools(env)...)
}
s.coreTools = append(s.coreTools, agentWebSearchTool(env))
s.coreTools = wrapFileToolsWithMemory(s.coreTools, s.args.Cwd)
s.coreTools = wrapContextAwareTools(s.coreTools, s.args.Cwd, discoverRules(s.args.Cwd))
s.deferredBase = agentLinearTools(env)
s.deferredBase = append(s.deferredBase, agentMemoryIndexTool(env))
Expand Down
6 changes: 2 additions & 4 deletions cmd/ask/aliases.go
Original file line number Diff line number Diff line change
Expand Up @@ -189,6 +189,7 @@ var (
agentWebSearchTool = tools.WebSearchTool
agentLoadMemoryTool = tools.LoadMemoryTool
agentPreloadMemoryTool = tools.PreloadMemoryTool
agentLoadArtifactsTool = tools.LoadArtifactsTool
)

const (
Expand Down Expand Up @@ -238,10 +239,6 @@ func agentMemorySystemBlock(cwd string) string {
return memory.SystemBlock(context.Background(), cwd)
}

func wrapFileToolsWithMemory(ts []tools.Tool, cwd string) []tools.Tool {
return tools.WrapFileToolsWithMemory(ts, cwd)
}

func agentMemoryIndexTool(env *agentToolEnv) tools.Tool {
return tools.MemoryIndexTool(env.Cwd, env.RequestApproval)
}
Expand Down Expand Up @@ -347,3 +344,4 @@ func buildAgentSystemPrompt(args ProviderSessionArgs) string {
GitStatusFn: agentGitStatus,
})
}
}
150 changes: 56 additions & 94 deletions cmd/ask/coordinator.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,15 @@ package main

import (
"context"
"errors"
"fmt"
"strings"
"sync"

tea "charm.land/bubbletea/v2"
"github.com/Cidan/ask/pkg/engine"
"github.com/Cidan/ask/pkg/providers"
"github.com/Cidan/ask/pkg/workflow"
adkmodel "google.golang.org/adk/v2/model"
"google.golang.org/adk/v2/tool"
)

// Coordinator manages the background execution of all in-process agent sessions
Expand Down Expand Up @@ -281,96 +282,6 @@ func (l tuiWorkflowListener) OnNote(tabID int, text string) {
}

// ExecuteStep implements workflow.StepExecutor for Coordinator.
func (c *Coordinator) ExecuteStep(ctx context.Context, cwd string, tabID int, step workflow.Step, prompt string, isFinal bool) (workflow.StepResult, error) {
prov := providerByID(step.Provider)
if prov == nil {
return workflow.StepResult{}, fmt.Errorf("provider not registered: %s", step.Provider)
}

args := ProviderSessionArgs{
Cwd: cwd,
TabID: tabID,
Model: step.Model,
Effort: "medium",
SkipAllPermissions: true,
InWorkflow: true,
IsWorkflowFinalStep: isFinal,
}

proc, ch, err := prov.StartSession(args)
if err != nil {
return workflow.StepResult{}, err
}

session, ok := proc.payload.(*agentSession)
if !ok {
return workflow.StepResult{}, errors.New("proc payload is not an agent session")
}
c.SetSession(tabID, session)

err = session.queueTurn(prompt)
if err != nil {
session.shutdown()
c.RemoveSession(tabID)
return workflow.StepResult{}, err
}

var stepResult string
var stepErr error
stepLoop:
for msg := range ch {
switch m := msg.(type) {
case assistantTextMsg:
stepResult += m.text
case providerDoneMsg:
if m.err != nil {
stepErr = m.err
} else if m.res.IsError {
stepErr = fmt.Errorf("step failed: %s", m.res.Result)
} else {
stepResult = m.res.Result
}
case turnCompleteMsg:
break stepLoop
}
}

session.shutdown()
c.RemoveSession(tabID)

if stepErr != nil {
return workflow.StepResult{}, stepErr
}

summary := ""
decision := ""
if session.env.PendingEndTurn != nil {
summary = session.env.PendingEndTurn.Summary
decision = session.env.PendingEndTurn.Decision
}
if summary == "" && strings.TrimSpace(stepResult) != "" {
firstLine := strings.TrimSpace(strings.Split(strings.TrimSpace(stepResult), "\n")[0])
if len(firstLine) > 200 {
firstLine = firstLine[:200] + "…"
}
summary = firstLine
}

var finishData *workflow.FinishData
if session.env.PendingFinishData != nil {
finishData = &workflow.FinishData{
Description: session.env.PendingFinishData.Description,
Artifacts: session.env.PendingFinishData.Artifacts,
}
}

return workflow.StepResult{
Output: stepResult,
Summary: summary,
Decision: decision,
FinishData: finishData,
}, nil
}

// RunWorkflow executes a workflow synchronously step by step in the background.
func (c *Coordinator) RunWorkflow(ctx context.Context, tabID int, def workflowDef, src workflowSource) (finalizedPlanReply, error) {
Expand Down Expand Up @@ -407,8 +318,59 @@ func (c *Coordinator) RunWorkflow(ctx context.Context, tabID int, def workflowDe
}

listener := tuiWorkflowListener{tabID: tabID}
runner := workflow.NewRunner(workflow.GlobalTracker(), c, listener)
runState, err := runner.Run(ctx, rootCwd, tabID, toPkgWorkflowDef(def), src)

cfg := workflow.WorkflowAgentConfig{
Def: toPkgWorkflowDef(def),
Source: src,
Cwd: rootCwd,
TabID: tabID,
ModelBuilder: func(ctx context.Context, step workflow.Step) (adkmodel.LLM, error) {
providerID := step.Provider
if providerID == "" {
providerID = "vertex"
}
spec, ok := providers.GetAgentProviderSpec(providerID)
if !ok || spec == nil {
return nil, fmt.Errorf("unknown provider %q", providerID)
}
config, _ := loadConfig()
modelID := providers.CanonicalVertexModelID(step.Model, spec.DefaultModel)
return engine.ModelBuilder(ctx, spec, toPkgConfig(config), modelID)
},
ToolsBuilder: func(ctx context.Context, step workflow.Step, isLoop bool) ([]tool.Tool, error) {
var agentTools []engine.Tool
if tf := engine.GetDefaultToolFactory(); tf != nil {
agentTools = tf(engine.ToolFactoryArgs{
Cwd: rootCwd,
TabID: tabID,
SkipPermissions: true,
AttachWebSearch: true,
})
}
return engine.AsADKTools(agentTools)
},
ToolsetsBuilder: func(ctx context.Context, step workflow.Step, isLoop bool) ([]tool.Toolset, error) {
var toolsets []tool.Toolset
if skillTS, err := engine.NewSkillToolset(ctx, rootCwd); err == nil && skillTS != nil {
toolsets = append(toolsets, skillTS)
}
return toolsets, nil
},
InstructionBuilder: func(step workflow.Step, isStart bool, isFinal bool, loopCtx *workflow.LoopPromptCtx, notesDir, prevNotesDir string) string {
pc := &workflow.StepPromptCtx{
Loop: loopCtx,
NotesDir: notesDir,
PrevNotesDir: prevNotesDir,
IsStartStep: isStart,
IsWorkflowFinalStep: isFinal,
}
return workflow.BuildStepPrompt(step, src, nil, pc)
},
SessionService: engine.NewFileSessionService("ask-workflow", rootCwd),
}

runner := workflow.NewRunner(workflow.GlobalTracker(), cfg)
runState, err := runner.Run(ctx, listener)
if err != nil {
return finalizedPlanReply{}, err
}
Expand Down
19 changes: 13 additions & 6 deletions pkg/engine/engine.go
Original file line number Diff line number Diff line change
Expand Up @@ -54,11 +54,12 @@ func (e *Engine) SystemPrompt(cwd string, inWorkflow bool) string {

// BuildWorkflowAgent constructs an ADK agent hierarchy (sequentialagent, loopagent, exitlooptool)
// for the given workflow definition using the engine's model and tool configuration.
func (e *Engine) BuildWorkflowAgent(ctx context.Context, cwd string, def workflow.Def, src workflow.Source) (agent.Agent, error) {
cfg := workflow.WorkflowAgentConfig{
func (e *Engine) BuildWorkflowAgentConfig(ctx context.Context, cwd string, tabID int, def workflow.Def, src workflow.Source) workflow.WorkflowAgentConfig {
return workflow.WorkflowAgentConfig{
Def: def,
Source: src,
Cwd: cwd,
TabID: tabID,
ModelBuilder: func(ctx context.Context, step workflow.Step) (model.LLM, error) {
providerID := step.Provider
if providerID == "" {
Expand Down Expand Up @@ -86,7 +87,7 @@ func (e *Engine) BuildWorkflowAgent(ctx context.Context, cwd string, def workflo
if tf := GetDefaultToolFactory(); tf != nil {
agentTools = tf(ToolFactoryArgs{
Cwd: cwd,
TabID: 0,
TabID: tabID,
SkipPermissions: true,
EventListener: e.opts.EventListener,
InteractionHandler: e.opts.InteractionHandler,
Expand All @@ -112,8 +113,13 @@ func (e *Engine) BuildWorkflowAgent(ctx context.Context, cwd string, def workflo
}
return workflow.BuildStepPrompt(step, src, nil, pc)
},
SessionService: NewFileSessionService("ask-workflow", cwd),
}
return workflow.BuildWorkflowAgent(ctx, cfg)
}

func (e *Engine) BuildWorkflowAgent(ctx context.Context, cwd string, def workflow.Def, src workflow.Source) (agent.Agent, error) {
cfg := e.BuildWorkflowAgentConfig(ctx, cwd, 0, def, src)
return workflow.CompileDefToADKWorkflow(ctx, cfg)
}

type engineWorkflowListener struct {
Expand Down Expand Up @@ -159,7 +165,8 @@ func (l engineWorkflowListener) OnNote(tabID int, text string) {

func (e *Engine) RunWorkflow(ctx context.Context, cwd string, tabID int, def workflow.Def, src workflow.Source) error {
listener := engineWorkflowListener{tabID: tabID, listener: e.opts.EventListener}
runner := workflow.NewRunner(workflow.GlobalTracker(), e.coordinator, listener)
_, err := runner.Run(ctx, cwd, tabID, def, src)
cfg := e.BuildWorkflowAgentConfig(ctx, cwd, tabID, def, src)
runner := workflow.NewRunner(workflow.GlobalTracker(), cfg)
_, err := runner.Run(ctx, listener)
return err
}
29 changes: 0 additions & 29 deletions pkg/engine/plugins.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,16 +17,6 @@ func DefaultPlugins() []*plugin.Plugin {
plugins = append(plugins, retryPlugin)
}

// 2. Function call modifier plugin for parameter injection when configured.
// Defaults to inactive so native tools using functiontool.New with ParametersJsonSchema
// are not corrupted by the plugin's decl.Parameters initialization.
if modPlugin, err := NewFunctionCallModifierPlugin(FunctionCallModifierOptions{
Predicate: func(toolName string) bool {
return false
},
}); err == nil && modPlugin != nil {
plugins = append(plugins, modPlugin)
}

return plugins
}
Expand All @@ -53,24 +43,5 @@ func NewRetryAndReflectPlugin(maxRetries int) (*plugin.Plugin, error) {
}

// FunctionCallModifierOptions defines configuration for the functioncallmodifier plugin.
type FunctionCallModifierOptions struct {
Predicate func(toolName string) bool
Args map[string]*genai.Schema
OverrideDescription func(originalDescription string) string
}

// NewFunctionCallModifierPlugin creates an ADK functioncallmodifier plugin with safety guards.
func NewFunctionCallModifierPlugin(opts FunctionCallModifierOptions) (*plugin.Plugin, error) {
pred := opts.Predicate
if pred == nil {
pred = func(toolName string) bool {
return len(opts.Args) > 0 || opts.OverrideDescription != nil
}
}
cfg := functioncallmodifier.FunctionCallModifierConfig{
Predicate: pred,
Args: opts.Args,
OverrideDescription: opts.OverrideDescription,
}
return functioncallmodifier.NewPlugin(cfg)
}
34 changes: 0 additions & 34 deletions pkg/engine/plugins_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -61,40 +61,6 @@ func TestNewRetryAndReflectPlugin(t *testing.T) {
})
}

func TestNewFunctionCallModifierPlugin(t *testing.T) {
t.Run("with nil predicate and args", func(t *testing.T) {
p, err := NewFunctionCallModifierPlugin(FunctionCallModifierOptions{})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if p == nil || p.Name() != "FunctionCallModifierPlugin" {
t.Fatalf("expected valid FunctionCallModifierPlugin, got %v", p)
}
})

t.Run("with custom schema args and description override", func(t *testing.T) {
p, err := NewFunctionCallModifierPlugin(FunctionCallModifierOptions{
Predicate: func(toolName string) bool {
return toolName == "read"
},
Args: map[string]*genai.Schema{
"description": {
Type: "STRING",
Description: "short phrase",
},
},
OverrideDescription: func(orig string) string {
return orig + " (modified)"
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if p == nil || p.Name() != "FunctionCallModifierPlugin" {
t.Fatalf("expected valid FunctionCallModifierPlugin, got %v", p)
}
})
}

func TestFunctionCallModifier_ActiveParameterInjection(t *testing.T) {
plugins := DefaultPlugins()
Expand Down
16 changes: 0 additions & 16 deletions pkg/engine/prompt.go
Original file line number Diff line number Diff line change
Expand Up @@ -496,22 +496,6 @@ func BuildSystemPrompt(opts PromptOptions) string {
}
}

if memory.IsOpen() {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
if mem := memory.SystemBlock(ctx, cwd); mem != "" {
b.WriteString("\n\n<project_memory>\n")
b.WriteString(mem)
b.WriteString("\n</project_memory>")
}
cancel()
}

if !opts.DisableSkillsPrompt {
if block := SkillsPromptBlock(DiscoverSkills(cwd)); block != "" {
b.WriteString("\n\n")
b.WriteString(block)
}
}
if block := SubagentsPromptBlock(DiscoverSubagents(cwd)); block != "" {
b.WriteString("\n\n")
b.WriteString(block)
Expand Down
Loading