feat(agent,sqlc,security): conversation/thread/mention ops + security additions
## Conversation Threading Operations Layer (pkg/agent/conversations/) Extracted from Budgetsmith ops/conversations. Adapted to use injected *sqlc.Queries instead of bootstrap.New(ctx). Full Store with: - List(ListParams) — paginated conversation listing, clamped [1,200] - Create(CreateParams) — new empty conversation - EditMessage(EditMessageParams) — update content + write revision record - ForkFromCheckpoint(ForkFromCheckpointParams) — branch at snapshot, seed up to 200 messages, write agent_conversation_forks lineage record - MergeAsLinkedContext(MergeAsLinkedContextParams) — new conversation with link records to two existing conversations; writes self-describing note - Ancestry(AncestryParams) — fork parent, fork children, merge links - LinksList / LinksRemove — manage conversation links directly - Graph(GraphParams) — depth-limited fork/link graph (max depth 10), upward+downward fork traversal with non-traversing link edge inclusion ## Thread Operations Layer (pkg/agent/threads/) Extracted from Budgetsmith ops/threads. Store with: - Create / List — thread CRUD within a conversation - AddMessage / ListMessages — append messages, return chronological order ## Mention Operations Layer (pkg/agent/mentions/) Store with Add / ListByConversation using agent_mentions schema (Kind/TargetID/Raw). ## SQLC additions (pkg/memory/sqlc/) - GetAgentMessageByID :one — fetch single message by primary key - UpdateAgentMessageContent :one — update content, RETURNING row (both needed by EditMessage to snapshot old content before update) ## Security additions - pkg/security/zkp.go + test: zero-knowledge proof helpers - pkg/security/securebus/socket_transport.go + test: Unix socket transport for out-of-process SecureBus communication - pkg/security/securebus/bus.go: minor fixes - pkg/itr/dag/executor.go + test: DAG executor refinements - pkg/tools/toolloop.go: tool execution loop improvements - cmd/picoclaw/main.go: integration wiring updates
This commit is contained in:
parent
32fcd87737
commit
395fef1157
15 changed files with 2605 additions and 35 deletions
|
|
@ -31,6 +31,7 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/devices"
|
||||
picofantasy "github.com/sipeed/picoclaw/pkg/fantasy"
|
||||
"github.com/sipeed/picoclaw/pkg/health"
|
||||
"github.com/sipeed/picoclaw/pkg/itr"
|
||||
"github.com/sipeed/picoclaw/pkg/heartbeat"
|
||||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
picomemory "github.com/sipeed/picoclaw/pkg/memory"
|
||||
|
|
@ -195,6 +196,10 @@ func main() {
|
|||
fmt.Printf("Unknown skills command: %s\n", subcommand)
|
||||
skillsHelp()
|
||||
}
|
||||
case "secret":
|
||||
secretCmd()
|
||||
case "daemon":
|
||||
daemonCmd()
|
||||
case "memory":
|
||||
memoryCmd()
|
||||
case "version", "--version", "-v":
|
||||
|
|
@ -219,6 +224,8 @@ func printHelp() {
|
|||
fmt.Println(" cron Manage scheduled tasks")
|
||||
fmt.Println(" migrate Migrate from OpenClaw to PicoClaw")
|
||||
fmt.Println(" memory Memory system management (db status, session migration)")
|
||||
fmt.Println(" secret Manage secrets (add, list, delete)")
|
||||
fmt.Println(" daemon Manage the picoclaw daemon (start, stop, status)")
|
||||
fmt.Println(" skills Manage skills (install, list, remove)")
|
||||
fmt.Println(" version Show version information")
|
||||
}
|
||||
|
|
@ -247,6 +254,26 @@ func onboard() {
|
|||
createWorkspaceTemplates(workspace)
|
||||
|
||||
fmt.Printf("%s picoclaw is ready!\n", logo)
|
||||
|
||||
fmt.Print("\nSet up encrypted secret storage? (y/n): ")
|
||||
var secretResponse string
|
||||
fmt.Scanln(&secretResponse)
|
||||
if secretResponse == "y" || secretResponse == "Y" {
|
||||
key, err := security.GenerateKey()
|
||||
if err != nil {
|
||||
fmt.Printf("Error generating key: %v\n", err)
|
||||
} else {
|
||||
encoded := fmt.Sprintf("%x", key)
|
||||
fmt.Println("\nGenerated master key (keep this safe!):")
|
||||
fmt.Println(" " + encoded)
|
||||
fmt.Println()
|
||||
fmt.Println("Add to your shell profile:")
|
||||
fmt.Println(" export PICOCLAW_MASTER_KEY=" + encoded)
|
||||
fmt.Println()
|
||||
fmt.Println("Then store secrets with: picoclaw secret add <name>")
|
||||
}
|
||||
}
|
||||
|
||||
fmt.Println("\nNext steps:")
|
||||
fmt.Println(" 1. Add your API key to", configPath)
|
||||
fmt.Println(" Get one at: https://openrouter.ai/keys")
|
||||
|
|
@ -1150,6 +1177,361 @@ func authStatusCmd() {
|
|||
}
|
||||
}
|
||||
|
||||
func secretCmd() {
|
||||
if len(os.Args) < 3 {
|
||||
secretHelp()
|
||||
return
|
||||
}
|
||||
|
||||
sub := os.Args[2]
|
||||
switch sub {
|
||||
case "init":
|
||||
secretInit()
|
||||
case "add":
|
||||
secretAdd()
|
||||
case "list":
|
||||
secretList()
|
||||
case "delete":
|
||||
secretDelete()
|
||||
case "--help", "-h":
|
||||
secretHelp()
|
||||
default:
|
||||
fmt.Printf("Unknown secret command: %s\n", sub)
|
||||
secretHelp()
|
||||
}
|
||||
}
|
||||
|
||||
func secretHelp() {
|
||||
fmt.Println("\nSecret management (encrypted at rest)")
|
||||
fmt.Println()
|
||||
fmt.Println("Usage: picoclaw secret <subcommand>")
|
||||
fmt.Println()
|
||||
fmt.Println("Subcommands:")
|
||||
fmt.Println(" init Generate a master key (stored in keyring or env)")
|
||||
fmt.Println(" add <name> Add or update a secret")
|
||||
fmt.Println(" list List stored secret names")
|
||||
fmt.Println(" delete <name> Remove a secret")
|
||||
fmt.Println()
|
||||
fmt.Println("Environment:")
|
||||
fmt.Println(" PICOCLAW_MASTER_KEY 32-byte key (hex or base64) for encryption")
|
||||
fmt.Println()
|
||||
fmt.Println("Examples:")
|
||||
fmt.Println(" picoclaw secret init")
|
||||
fmt.Println(" picoclaw secret add github_token")
|
||||
fmt.Println(" picoclaw secret list")
|
||||
fmt.Println(" picoclaw secret delete github_token")
|
||||
}
|
||||
|
||||
func secretStorePath() string {
|
||||
home, _ := os.UserHomeDir()
|
||||
return filepath.Join(home, ".picoclaw", "secrets.json")
|
||||
}
|
||||
|
||||
func loadSecretStore() (*security.SecretStore, error) {
|
||||
var keyring security.KeyringProvider
|
||||
if mk := os.Getenv("PICOCLAW_MASTER_KEY"); mk != "" {
|
||||
keyring = security.NewEnvKeyring("PICOCLAW_MASTER_KEY")
|
||||
} else {
|
||||
keyring = security.NewNoopKeyring(nil)
|
||||
}
|
||||
return security.NewSecretStore(secretStorePath(), keyring)
|
||||
}
|
||||
|
||||
func secretInit() {
|
||||
key, err := security.GenerateKey()
|
||||
if err != nil {
|
||||
fmt.Printf("Error generating key: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
encoded := fmt.Sprintf("%x", key)
|
||||
fmt.Println("Generated master key (keep this safe!):")
|
||||
fmt.Println()
|
||||
fmt.Println(" " + encoded)
|
||||
fmt.Println()
|
||||
fmt.Println("Set it as an environment variable:")
|
||||
fmt.Println(" export PICOCLAW_MASTER_KEY=" + encoded)
|
||||
fmt.Println()
|
||||
fmt.Println("Or add to your shell profile (~/.bashrc, ~/.zshrc).")
|
||||
}
|
||||
|
||||
func secretAdd() {
|
||||
if len(os.Args) < 4 {
|
||||
fmt.Println("Usage: picoclaw secret add <name>")
|
||||
return
|
||||
}
|
||||
name := os.Args[3]
|
||||
|
||||
ss, err := loadSecretStore()
|
||||
if err != nil {
|
||||
fmt.Printf("Error loading secret store: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("Enter value for %q (input hidden): ", name)
|
||||
reader := bufio.NewReader(os.Stdin)
|
||||
value, err := reader.ReadString('\n')
|
||||
if err != nil {
|
||||
fmt.Printf("Error reading input: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
value = strings.TrimSpace(value)
|
||||
|
||||
if value == "" {
|
||||
fmt.Println("Error: secret value cannot be empty")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
if err := ss.Set(name, []byte(value)); err != nil {
|
||||
fmt.Printf("Error storing secret: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("Secret %q stored (%s)\n", name, secretStorePath())
|
||||
}
|
||||
|
||||
func secretList() {
|
||||
ss, err := loadSecretStore()
|
||||
if err != nil {
|
||||
fmt.Printf("Error loading secret store: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
names := ss.List()
|
||||
if len(names) == 0 {
|
||||
fmt.Println("No secrets stored.")
|
||||
fmt.Println("Add one with: picoclaw secret add <name>")
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Printf("\nStored secrets (%d):\n", len(names))
|
||||
for _, name := range names {
|
||||
fmt.Printf(" - %s\n", name)
|
||||
}
|
||||
}
|
||||
|
||||
func secretDelete() {
|
||||
if len(os.Args) < 4 {
|
||||
fmt.Println("Usage: picoclaw secret delete <name>")
|
||||
return
|
||||
}
|
||||
name := os.Args[3]
|
||||
|
||||
ss, err := loadSecretStore()
|
||||
if err != nil {
|
||||
fmt.Printf("Error loading secret store: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
if !ss.Has(name) {
|
||||
fmt.Printf("Secret %q not found\n", name)
|
||||
return
|
||||
}
|
||||
|
||||
if err := ss.Delete(name); err != nil {
|
||||
fmt.Printf("Error deleting secret: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("Secret %q deleted\n", name)
|
||||
}
|
||||
|
||||
func daemonSocketPath() string {
|
||||
home, _ := os.UserHomeDir()
|
||||
return filepath.Join(home, ".picoclaw", "daemon.sock")
|
||||
}
|
||||
|
||||
func daemonPIDPath() string {
|
||||
home, _ := os.UserHomeDir()
|
||||
return filepath.Join(home, ".picoclaw", "daemon.pid")
|
||||
}
|
||||
|
||||
func daemonCmd() {
|
||||
if len(os.Args) < 3 {
|
||||
daemonHelp()
|
||||
return
|
||||
}
|
||||
|
||||
sub := os.Args[2]
|
||||
switch sub {
|
||||
case "start":
|
||||
daemonStart()
|
||||
case "stop":
|
||||
daemonStop()
|
||||
case "status":
|
||||
daemonStatus()
|
||||
case "--help", "-h":
|
||||
daemonHelp()
|
||||
default:
|
||||
fmt.Printf("Unknown daemon command: %s\n", sub)
|
||||
daemonHelp()
|
||||
}
|
||||
}
|
||||
|
||||
func daemonHelp() {
|
||||
fmt.Println("\nDaemon mode (Unix socket transport)")
|
||||
fmt.Println()
|
||||
fmt.Println("Usage: picoclaw daemon <subcommand>")
|
||||
fmt.Println()
|
||||
fmt.Println("Subcommands:")
|
||||
fmt.Println(" start Start the daemon (foreground)")
|
||||
fmt.Println(" stop Stop a running daemon")
|
||||
fmt.Println(" status Check daemon status")
|
||||
fmt.Println()
|
||||
fmt.Println("The daemon listens on ~/.picoclaw/daemon.sock and provides")
|
||||
fmt.Println("tool execution services via the SecureBus.")
|
||||
}
|
||||
|
||||
func daemonStart() {
|
||||
sockPath := daemonSocketPath()
|
||||
pidPath := daemonPIDPath()
|
||||
|
||||
if data, err := os.ReadFile(pidPath); err == nil {
|
||||
fmt.Printf("Daemon PID file exists (%s): %s\n", pidPath, strings.TrimSpace(string(data)))
|
||||
fmt.Println("If the daemon is not running, remove the PID file and try again:")
|
||||
fmt.Printf(" rm %s\n", pidPath)
|
||||
return
|
||||
}
|
||||
|
||||
server, err := securebus.NewSocketTransportServer(sockPath)
|
||||
if err != nil {
|
||||
fmt.Printf("Error creating socket: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer server.Close()
|
||||
|
||||
pid := os.Getpid()
|
||||
home, _ := os.UserHomeDir()
|
||||
_ = os.MkdirAll(filepath.Join(home, ".picoclaw"), 0700)
|
||||
_ = os.WriteFile(pidPath, []byte(fmt.Sprintf("%d", pid)), 0600)
|
||||
defer os.Remove(pidPath)
|
||||
|
||||
cfg, err := loadConfig()
|
||||
if err != nil {
|
||||
fmt.Printf("Error loading config: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
ss, err := loadSecretStore()
|
||||
if err != nil {
|
||||
logger.WarnCF("daemon", "failed to load secret store", map[string]interface{}{"error": err.Error()})
|
||||
ss = nil
|
||||
}
|
||||
|
||||
registry := tools.NewToolRegistry()
|
||||
workspace := cfg.WorkspacePath()
|
||||
restrict := cfg.Agents.Defaults.RestrictToWorkspace
|
||||
registry.Register(tools.NewExecTool(workspace, restrict))
|
||||
registry.Register(tools.NewReadFileTool(workspace, restrict))
|
||||
registry.Register(tools.NewWriteFileTool(workspace, restrict))
|
||||
registry.Register(tools.NewListDirTool(workspace, restrict))
|
||||
registry.Register(tools.NewEditFileTool(workspace, restrict))
|
||||
|
||||
capLookup := func(name string) (tools.ToolCapabilities, bool) {
|
||||
t, ok := registry.Get(name)
|
||||
if !ok {
|
||||
return tools.ZeroCapabilities(), false
|
||||
}
|
||||
return tools.ExtractCapabilities(t), true
|
||||
}
|
||||
executor := func(ctx context.Context, name string, args map[string]interface{}) *tools.ToolResult {
|
||||
return registry.Execute(ctx, name, args)
|
||||
}
|
||||
|
||||
busCfg := securebus.DefaultBusConfig()
|
||||
secureBus := securebus.New(busCfg, ss, capLookup, executor)
|
||||
defer secureBus.Close()
|
||||
|
||||
fmt.Printf("picoclaw daemon started (pid=%d, socket=%s)\n", pid, sockPath)
|
||||
fmt.Printf(" workspace: %s\n", workspace)
|
||||
fmt.Printf(" tools: %d registered\n", len(registry.List()))
|
||||
fmt.Println("Press Ctrl+C to stop.")
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt)
|
||||
defer stop()
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
errCh <- server.Serve(func(srvCtx context.Context, req itr.ToolRequest) itr.ToolResponse {
|
||||
return secureBus.Execute(srvCtx, req)
|
||||
})
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
fmt.Println("\nShutting down daemon...")
|
||||
case err := <-errCh:
|
||||
if err != nil {
|
||||
fmt.Printf("Daemon error: %v\n", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func daemonStop() {
|
||||
pidPath := daemonPIDPath()
|
||||
|
||||
data, err := os.ReadFile(pidPath)
|
||||
if err != nil {
|
||||
fmt.Println("No daemon PID file found. Is the daemon running?")
|
||||
return
|
||||
}
|
||||
|
||||
pidStr := strings.TrimSpace(string(data))
|
||||
var pid int
|
||||
if _, err := fmt.Sscanf(pidStr, "%d", &pid); err != nil {
|
||||
fmt.Printf("Invalid PID in %s: %s\n", pidPath, pidStr)
|
||||
return
|
||||
}
|
||||
|
||||
proc, err := os.FindProcess(pid)
|
||||
if err != nil {
|
||||
fmt.Printf("Could not find process %d: %v\n", pid, err)
|
||||
_ = os.Remove(pidPath)
|
||||
return
|
||||
}
|
||||
|
||||
if err := proc.Signal(os.Interrupt); err != nil {
|
||||
fmt.Printf("Could not signal process %d: %v\n", pid, err)
|
||||
fmt.Println("Removing stale PID file.")
|
||||
_ = os.Remove(pidPath)
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Printf("Sent interrupt to daemon (pid=%d)\n", pid)
|
||||
}
|
||||
|
||||
func daemonStatus() {
|
||||
pidPath := daemonPIDPath()
|
||||
sockPath := daemonSocketPath()
|
||||
|
||||
data, err := os.ReadFile(pidPath)
|
||||
if err != nil {
|
||||
fmt.Println("Daemon: not running (no PID file)")
|
||||
return
|
||||
}
|
||||
|
||||
pidStr := strings.TrimSpace(string(data))
|
||||
fmt.Printf("Daemon PID: %s\n", pidStr)
|
||||
|
||||
if _, err := os.Stat(sockPath); err == nil {
|
||||
fmt.Printf("Socket: %s (exists)\n", sockPath)
|
||||
} else {
|
||||
fmt.Printf("Socket: %s (missing)\n", sockPath)
|
||||
}
|
||||
|
||||
var pid int
|
||||
if _, err := fmt.Sscanf(pidStr, "%d", &pid); err == nil {
|
||||
proc, err := os.FindProcess(pid)
|
||||
if err == nil {
|
||||
if err := proc.Signal(nil); err == nil {
|
||||
fmt.Println("Status: running")
|
||||
} else {
|
||||
fmt.Println("Status: stale PID file (process not found)")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func getConfigPath() string {
|
||||
home, _ := os.UserHomeDir()
|
||||
return filepath.Join(home, ".picoclaw", "config.json")
|
||||
|
|
|
|||
562
pkg/agent/conversations/store.go
Normal file
562
pkg/agent/conversations/store.go
Normal file
|
|
@ -0,0 +1,562 @@
|
|||
// Package conversations provides the business-logic operations layer for
|
||||
// conversation management (create, list, edit, fork, merge, ancestry, graph).
|
||||
// It sits above the raw SQLC-generated queries and below any HTTP/CLI handler.
|
||||
package conversations
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/ids"
|
||||
sqlc "github.com/sipeed/picoclaw/pkg/memory/sqlc"
|
||||
"github.com/sipeed/picoclaw/pkg/pcerrors"
|
||||
)
|
||||
|
||||
// Store wraps the SQLC queries for conversation operations.
|
||||
// Construct with New.
|
||||
type Store struct {
|
||||
q *sqlc.Queries
|
||||
}
|
||||
|
||||
// New returns a Store backed by the provided SQLC queries handle.
|
||||
func New(q *sqlc.Queries) *Store {
|
||||
return &Store{q: q}
|
||||
}
|
||||
|
||||
// ─── List ────────────────────────────────────────────────────────────────────
|
||||
|
||||
// ListParams configures the List operation.
|
||||
type ListParams struct {
|
||||
// Limit is clamped to [1, 200]; defaults to 20.
|
||||
Limit int
|
||||
}
|
||||
|
||||
// List returns up to Limit conversations ordered by creation time (newest first).
|
||||
func (s *Store) List(ctx context.Context, p ListParams) ([]sqlc.AgentConversation, error) {
|
||||
limit := int64(p.Limit)
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 200 {
|
||||
limit = 200
|
||||
}
|
||||
return s.q.ListAgentConversations(ctx, sqlc.ListAgentConversationsParams{Limit: limit})
|
||||
}
|
||||
|
||||
// ─── Create ──────────────────────────────────────────────────────────────────
|
||||
|
||||
// CreateParams configures the Create operation.
|
||||
type CreateParams struct {
|
||||
Title *string
|
||||
}
|
||||
|
||||
// Create creates a new, empty conversation.
|
||||
func (s *Store) Create(ctx context.Context, p CreateParams) (sqlc.AgentConversation, error) {
|
||||
return s.q.CreateAgentConversation(ctx, sqlc.CreateAgentConversationParams{
|
||||
ID: ids.New(),
|
||||
Title: p.Title,
|
||||
})
|
||||
}
|
||||
|
||||
// ─── EditMessage ──────────────────────────────────────────────────────────────
|
||||
|
||||
// EditMessageParams configures the EditMessage operation.
|
||||
type EditMessageParams struct {
|
||||
// MessageID is the UUID (string) of the message to edit.
|
||||
MessageID string
|
||||
// NewText is the replacement content. Must be non-empty.
|
||||
NewText string
|
||||
// Editor is recorded in the revision history. Defaults to "user".
|
||||
Editor string
|
||||
// Metadata is serialised to JSON and stored on the revision row.
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
// EditMessage updates a message's content and records the old value in the
|
||||
// revision table for audit/undo purposes.
|
||||
func (s *Store) EditMessage(ctx context.Context, p EditMessageParams) (sqlc.AgentMessage, error) {
|
||||
msgIDStr := strings.TrimSpace(p.MessageID)
|
||||
if msgIDStr == "" {
|
||||
return sqlc.AgentMessage{}, pcerrors.New(pcerrors.CodeInvalidArgument, "message_id is required")
|
||||
}
|
||||
msgID, err := ids.Parse(msgIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentMessage{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse message_id %q", msgIDStr)
|
||||
}
|
||||
|
||||
newText := strings.TrimSpace(p.NewText)
|
||||
if newText == "" {
|
||||
return sqlc.AgentMessage{}, pcerrors.New(pcerrors.CodeInvalidArgument, "new content is empty")
|
||||
}
|
||||
|
||||
editor := strings.TrimSpace(p.Editor)
|
||||
if editor == "" {
|
||||
editor = "user"
|
||||
}
|
||||
|
||||
prev, err := s.q.GetAgentMessageByID(ctx, sqlc.GetAgentMessageByIDParams{ID: msgID})
|
||||
if err != nil {
|
||||
return sqlc.AgentMessage{}, err
|
||||
}
|
||||
|
||||
metaJSON, _ := json.Marshal(p.Metadata)
|
||||
|
||||
// Best-effort revision record — a failure here does not abort the edit.
|
||||
_, _ = s.q.AddAgentMessageRevision(ctx, sqlc.AddAgentMessageRevisionParams{
|
||||
ID: ids.New(),
|
||||
MessageID: msgID,
|
||||
Editor: editor,
|
||||
OldContent: prev.Content,
|
||||
NewContent: newText,
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
|
||||
updated, err := s.q.UpdateAgentMessageContent(ctx, sqlc.UpdateAgentMessageContentParams{
|
||||
Content: newText,
|
||||
ID: msgID,
|
||||
})
|
||||
if err != nil {
|
||||
return sqlc.AgentMessage{}, err
|
||||
}
|
||||
return updated, nil
|
||||
}
|
||||
|
||||
// ─── Fork ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
// ForkFromCheckpointParams configures the ForkFromCheckpoint operation.
|
||||
type ForkFromCheckpointParams struct {
|
||||
// FromConversationID is the UUID (string) of the source conversation.
|
||||
FromConversationID string
|
||||
// CheckpointName identifies the snapshot to fork from.
|
||||
CheckpointName string
|
||||
// Title is the optional title for the newly forked conversation.
|
||||
Title *string
|
||||
}
|
||||
|
||||
// ForkFromCheckpoint creates a new conversation branched off at the state
|
||||
// captured by the named checkpoint. Messages up to 200 are seeded into the
|
||||
// new conversation; a fork-lineage record is written to agent_conversation_forks.
|
||||
func (s *Store) ForkFromCheckpoint(ctx context.Context, p ForkFromCheckpointParams) (sqlc.AgentConversation, error) {
|
||||
fromIDStr := strings.TrimSpace(p.FromConversationID)
|
||||
if fromIDStr == "" {
|
||||
return sqlc.AgentConversation{}, pcerrors.New(pcerrors.CodeInvalidArgument, "from_conversation_id is required")
|
||||
}
|
||||
fromID, err := ids.Parse(fromIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse from_conversation_id %q", fromIDStr)
|
||||
}
|
||||
|
||||
cpName := strings.TrimSpace(p.CheckpointName)
|
||||
if cpName == "" {
|
||||
return sqlc.AgentConversation{}, pcerrors.New(pcerrors.CodeInvalidArgument, "checkpoint_name is required")
|
||||
}
|
||||
|
||||
cp, err := s.q.GetAgentCheckpointByConversationIDAndName(ctx,
|
||||
sqlc.GetAgentCheckpointByConversationIDAndNameParams{
|
||||
ConversationID: fromID,
|
||||
Name: cpName,
|
||||
})
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, err
|
||||
}
|
||||
|
||||
runState, err := s.q.GetAgentRunStateByID(ctx, sqlc.GetAgentRunStateByIDParams{ID: cp.RunStateID})
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, err
|
||||
}
|
||||
|
||||
type msgSnapshot struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
type snapshot struct {
|
||||
Messages []msgSnapshot `json:"messages"`
|
||||
}
|
||||
var snap snapshot
|
||||
_ = json.Unmarshal(runState.SnapshotJson, &snap)
|
||||
|
||||
conv, err := s.q.CreateAgentConversation(ctx, sqlc.CreateAgentConversationParams{
|
||||
ID: ids.New(),
|
||||
Title: p.Title,
|
||||
})
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, err
|
||||
}
|
||||
|
||||
forkMeta := map[string]any{
|
||||
"checkpoint_name": cpName,
|
||||
"run_state_id": cp.RunStateID.String(),
|
||||
}
|
||||
forkMetaJSON, _ := json.Marshal(forkMeta)
|
||||
|
||||
_, _ = s.q.CreateAgentConversationFork(ctx, sqlc.CreateAgentConversationForkParams{
|
||||
ID: ids.New(),
|
||||
ParentConversationID: fromID,
|
||||
ChildConversationID: conv.ID,
|
||||
CheckpointID: cp.ID,
|
||||
MetadataJson: forkMetaJSON,
|
||||
})
|
||||
|
||||
// Seed messages from snapshot — cap to 200 to prevent pathological snapshots.
|
||||
msgs := snap.Messages
|
||||
if len(msgs) > 200 {
|
||||
msgs = msgs[len(msgs)-200:]
|
||||
}
|
||||
|
||||
seedMeta := map[string]any{
|
||||
"seeded_from_conversation_id": fromID.String(),
|
||||
"checkpoint_name": cpName,
|
||||
"run_state_id": cp.RunStateID.String(),
|
||||
}
|
||||
seedMetaJSON, _ := json.Marshal(seedMeta)
|
||||
|
||||
for _, m := range msgs {
|
||||
if m.Role != "user" && m.Role != "assistant" {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(m.Content) == "" {
|
||||
continue
|
||||
}
|
||||
_, _ = s.q.AddAgentMessage(ctx, sqlc.AddAgentMessageParams{
|
||||
ID: ids.New(),
|
||||
ConversationID: conv.ID,
|
||||
Role: m.Role,
|
||||
Content: m.Content,
|
||||
MetadataJson: seedMetaJSON,
|
||||
})
|
||||
}
|
||||
|
||||
return conv, nil
|
||||
}
|
||||
|
||||
// ─── Merge ────────────────────────────────────────────────────────────────────
|
||||
|
||||
// MergeAsLinkedContextParams configures the MergeAsLinkedContext operation.
|
||||
type MergeAsLinkedContextParams struct {
|
||||
// BaseConversationID and OtherConversationID must differ.
|
||||
BaseConversationID string
|
||||
OtherConversationID string
|
||||
// Title is the optional title for the merged conversation.
|
||||
Title *string
|
||||
}
|
||||
|
||||
// MergeAsLinkedContext creates a new conversation that carries link records
|
||||
// to both source conversations. No messages are copied; linked context is
|
||||
// injected at runtime as a compact view. A system note is written so the
|
||||
// conversation is self-describing.
|
||||
func (s *Store) MergeAsLinkedContext(ctx context.Context, p MergeAsLinkedContextParams) (sqlc.AgentConversation, error) {
|
||||
baseIDStr := strings.TrimSpace(p.BaseConversationID)
|
||||
if baseIDStr == "" {
|
||||
return sqlc.AgentConversation{}, pcerrors.New(pcerrors.CodeInvalidArgument, "base_conversation_id is required")
|
||||
}
|
||||
baseID, err := ids.Parse(baseIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse base_conversation_id %q", baseIDStr)
|
||||
}
|
||||
|
||||
otherIDStr := strings.TrimSpace(p.OtherConversationID)
|
||||
if otherIDStr == "" {
|
||||
return sqlc.AgentConversation{}, pcerrors.New(pcerrors.CodeInvalidArgument, "other_conversation_id is required")
|
||||
}
|
||||
otherID, err := ids.Parse(otherIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse other_conversation_id %q", otherIDStr)
|
||||
}
|
||||
|
||||
if baseID == otherID {
|
||||
return sqlc.AgentConversation{}, pcerrors.New(pcerrors.CodeInvalidArgument, "base and other conversations must differ")
|
||||
}
|
||||
|
||||
conv, err := s.q.CreateAgentConversation(ctx, sqlc.CreateAgentConversationParams{
|
||||
ID: ids.New(),
|
||||
Title: p.Title,
|
||||
})
|
||||
if err != nil {
|
||||
return sqlc.AgentConversation{}, err
|
||||
}
|
||||
|
||||
meta := map[string]any{"source": "merge_as_linked_context"}
|
||||
metaJSON, _ := json.Marshal(meta)
|
||||
|
||||
_, _ = s.q.CreateAgentConversationLink(ctx, sqlc.CreateAgentConversationLinkParams{
|
||||
ID: ids.New(),
|
||||
ConversationID: conv.ID,
|
||||
LinkedConversationID: baseID,
|
||||
Kind: "merge",
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
_, _ = s.q.CreateAgentConversationLink(ctx, sqlc.CreateAgentConversationLinkParams{
|
||||
ID: ids.New(),
|
||||
ConversationID: conv.ID,
|
||||
LinkedConversationID: otherID,
|
||||
Kind: "merge",
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
|
||||
note := "This conversation was created by merging as linked context.\n" +
|
||||
"Linked conversations:\n" +
|
||||
"- @conv:" + baseID.String() + "\n" +
|
||||
"- @conv:" + otherID.String() + "\n\n" +
|
||||
"These links are user-attached context; the agent will be shown compact context from them at the start of each turn.\n"
|
||||
|
||||
_, _ = s.q.AddAgentMessage(ctx, sqlc.AddAgentMessageParams{
|
||||
ID: ids.New(),
|
||||
ConversationID: conv.ID,
|
||||
Role: "system",
|
||||
Content: note,
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
|
||||
return conv, nil
|
||||
}
|
||||
|
||||
// ─── Ancestry ────────────────────────────────────────────────────────────────
|
||||
|
||||
// AncestryParams configures the Ancestry operation.
|
||||
type AncestryParams struct {
|
||||
ConversationID string
|
||||
}
|
||||
|
||||
// AncestryResult is the structured result of an Ancestry query.
|
||||
type AncestryResult struct {
|
||||
Conversation sqlc.AgentConversation `json:"conversation"`
|
||||
ForkParent *sqlc.AgentConversationFork `json:"fork_parent,omitempty"`
|
||||
ForkChildren []sqlc.AgentConversationFork `json:"fork_children"`
|
||||
Links []sqlc.AgentConversationLink `json:"links"`
|
||||
}
|
||||
|
||||
// Ancestry returns the fork lineage and merge links for a conversation.
|
||||
func (s *Store) Ancestry(ctx context.Context, p AncestryParams) (AncestryResult, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return AncestryResult{}, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return AncestryResult{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
|
||||
conv, err := s.q.GetAgentConversation(ctx, sqlc.GetAgentConversationParams{ID: convID})
|
||||
if err != nil {
|
||||
return AncestryResult{}, err
|
||||
}
|
||||
|
||||
var parent *sqlc.AgentConversationFork
|
||||
if fp, err := s.q.GetAgentConversationForkByChildConversationID(ctx,
|
||||
sqlc.GetAgentConversationForkByChildConversationIDParams{ChildConversationID: convID}); err == nil {
|
||||
parent = &fp
|
||||
}
|
||||
|
||||
children, _ := s.q.ListAgentConversationForksByParentConversationID(ctx,
|
||||
sqlc.ListAgentConversationForksByParentConversationIDParams{ParentConversationID: convID})
|
||||
links, _ := s.q.ListAgentConversationLinksByConversationID(ctx,
|
||||
sqlc.ListAgentConversationLinksByConversationIDParams{ConversationID: convID})
|
||||
|
||||
return AncestryResult{
|
||||
Conversation: conv,
|
||||
ForkParent: parent,
|
||||
ForkChildren: children,
|
||||
Links: links,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ─── Links ────────────────────────────────────────────────────────────────────
|
||||
|
||||
// LinksListParams configures the LinksList operation.
|
||||
type LinksListParams struct {
|
||||
ConversationID string
|
||||
}
|
||||
|
||||
// LinksList returns all conversation links for the given conversation.
|
||||
func (s *Store) LinksList(ctx context.Context, p LinksListParams) ([]sqlc.AgentConversationLink, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return nil, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return nil, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
return s.q.ListAgentConversationLinksByConversationID(ctx,
|
||||
sqlc.ListAgentConversationLinksByConversationIDParams{ConversationID: convID})
|
||||
}
|
||||
|
||||
// LinksRemoveParams configures the LinksRemove operation.
|
||||
type LinksRemoveParams struct {
|
||||
ConversationID string
|
||||
LinkedConversationID string
|
||||
// Kind defaults to "merge".
|
||||
Kind string
|
||||
}
|
||||
|
||||
// LinksRemove deletes a specific conversation link.
|
||||
func (s *Store) LinksRemove(ctx context.Context, p LinksRemoveParams) error {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
|
||||
linkedIDStr := strings.TrimSpace(p.LinkedConversationID)
|
||||
if linkedIDStr == "" {
|
||||
return pcerrors.New(pcerrors.CodeInvalidArgument, "linked_conversation_id is required")
|
||||
}
|
||||
linkedID, err := ids.Parse(linkedIDStr)
|
||||
if err != nil {
|
||||
return pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse linked_conversation_id %q", linkedIDStr)
|
||||
}
|
||||
|
||||
kind := strings.TrimSpace(p.Kind)
|
||||
if kind == "" {
|
||||
kind = "merge"
|
||||
}
|
||||
|
||||
return s.q.DeleteAgentConversationLink(ctx, sqlc.DeleteAgentConversationLinkParams{
|
||||
ConversationID: convID,
|
||||
LinkedConversationID: linkedID,
|
||||
Kind: kind,
|
||||
})
|
||||
}
|
||||
|
||||
// ─── Graph ────────────────────────────────────────────────────────────────────
|
||||
|
||||
// GraphParams configures the Graph operation.
|
||||
type GraphParams struct {
|
||||
ConversationID string
|
||||
// Depth controls how many hops of fork lineage to traverse. Clamped [0,10].
|
||||
Depth int
|
||||
}
|
||||
|
||||
// GraphNode is a single conversation node in the fork/link graph.
|
||||
type GraphNode struct {
|
||||
ID string `json:"id"`
|
||||
Title *string `json:"title,omitempty"`
|
||||
}
|
||||
|
||||
// GraphEdge is a directed edge between two conversation nodes.
|
||||
type GraphEdge struct {
|
||||
// Type is "fork" or "link".
|
||||
Type string `json:"type"`
|
||||
From string `json:"from"`
|
||||
To string `json:"to"`
|
||||
// CheckpointID is set on fork edges.
|
||||
CheckpointID *string `json:"checkpoint_id,omitempty"`
|
||||
// Kind is set on link edges (e.g. "merge").
|
||||
Kind *string `json:"kind,omitempty"`
|
||||
}
|
||||
|
||||
// GraphResult is the full fork/link graph rooted at a conversation.
|
||||
type GraphResult struct {
|
||||
RootID string `json:"root_id"`
|
||||
Nodes []GraphNode `json:"nodes"`
|
||||
Edges []GraphEdge `json:"edges"`
|
||||
}
|
||||
|
||||
// Graph builds a depth-limited fork/link graph centred on a conversation.
|
||||
// - Fork edges are traversed up (to parent) and down (to children) up to Depth hops.
|
||||
// - Link edges are included for visited nodes but NOT traversed.
|
||||
func (s *Store) Graph(ctx context.Context, p GraphParams) (GraphResult, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return GraphResult{}, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
rootID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return GraphResult{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
|
||||
depth := p.Depth
|
||||
if depth < 0 {
|
||||
depth = 0
|
||||
}
|
||||
if depth > 10 {
|
||||
depth = 10
|
||||
}
|
||||
|
||||
nodes := map[string]GraphNode{}
|
||||
edges := map[string]GraphEdge{}
|
||||
visited := map[string]bool{}
|
||||
|
||||
var visit func(id ids.UUID, remaining int) error
|
||||
visit = func(id ids.UUID, remaining int) error {
|
||||
key := id.String()
|
||||
if visited[key] {
|
||||
return nil
|
||||
}
|
||||
visited[key] = true
|
||||
|
||||
conv, err := s.q.GetAgentConversation(ctx, sqlc.GetAgentConversationParams{ID: id})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
nodes[key] = GraphNode{ID: conv.ID.String(), Title: conv.Title}
|
||||
|
||||
// Include link edges but do not traverse them.
|
||||
if lks, err := s.q.ListAgentConversationLinksByConversationID(ctx,
|
||||
sqlc.ListAgentConversationLinksByConversationIDParams{ConversationID: id}); err == nil {
|
||||
for _, l := range lks {
|
||||
from := id.String()
|
||||
to := l.LinkedConversationID.String()
|
||||
kind := l.Kind
|
||||
edges["link:"+from+":"+to+":"+kind] = GraphEdge{
|
||||
Type: "link", From: from, To: to, Kind: &kind,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fork parent (upward traversal).
|
||||
if fp, err := s.q.GetAgentConversationForkByChildConversationID(ctx,
|
||||
sqlc.GetAgentConversationForkByChildConversationIDParams{ChildConversationID: id}); err == nil {
|
||||
from := fp.ParentConversationID.String()
|
||||
to := fp.ChildConversationID.String()
|
||||
cpID := fp.CheckpointID.String()
|
||||
edges["fork:"+from+":"+to] = GraphEdge{
|
||||
Type: "fork", From: from, To: to, CheckpointID: &cpID,
|
||||
}
|
||||
if remaining > 0 {
|
||||
_ = visit(fp.ParentConversationID, remaining-1)
|
||||
}
|
||||
}
|
||||
|
||||
// Fork children (downward traversal).
|
||||
if children, err := s.q.ListAgentConversationForksByParentConversationID(ctx,
|
||||
sqlc.ListAgentConversationForksByParentConversationIDParams{ParentConversationID: id}); err == nil {
|
||||
for _, c := range children {
|
||||
from := c.ParentConversationID.String()
|
||||
to := c.ChildConversationID.String()
|
||||
cpID := c.CheckpointID.String()
|
||||
edges["fork:"+from+":"+to] = GraphEdge{
|
||||
Type: "fork", From: from, To: to, CheckpointID: &cpID,
|
||||
}
|
||||
if remaining > 0 {
|
||||
_ = visit(c.ChildConversationID, remaining-1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := visit(rootID, depth); err != nil {
|
||||
return GraphResult{}, err
|
||||
}
|
||||
|
||||
outNodes := make([]GraphNode, 0, len(nodes))
|
||||
for _, n := range nodes {
|
||||
outNodes = append(outNodes, n)
|
||||
}
|
||||
outEdges := make([]GraphEdge, 0, len(edges))
|
||||
for _, e := range edges {
|
||||
outEdges = append(outEdges, e)
|
||||
}
|
||||
|
||||
return GraphResult{
|
||||
RootID: rootID.String(),
|
||||
Nodes: outNodes,
|
||||
Edges: outEdges,
|
||||
}, nil
|
||||
}
|
||||
111
pkg/agent/mentions/store.go
Normal file
111
pkg/agent/mentions/store.go
Normal file
|
|
@ -0,0 +1,111 @@
|
|||
// Package mentions provides the operations layer for agent mentions —
|
||||
// @-style references to conversations, threads, or documents embedded in
|
||||
// messages. Mentions are stored in agent_mentions for fast cross-entity lookup.
|
||||
package mentions
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/ids"
|
||||
sqlc "github.com/sipeed/picoclaw/pkg/memory/sqlc"
|
||||
"github.com/sipeed/picoclaw/pkg/pcerrors"
|
||||
)
|
||||
|
||||
// Store wraps the SQLC queries for mention operations.
|
||||
// Construct with New.
|
||||
type Store struct {
|
||||
q *sqlc.Queries
|
||||
}
|
||||
|
||||
// New returns a Store backed by the provided SQLC queries handle.
|
||||
func New(q *sqlc.Queries) *Store {
|
||||
return &Store{q: q}
|
||||
}
|
||||
|
||||
// ─── Add ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
// AddParams configures the Add operation.
|
||||
type AddParams struct {
|
||||
// ConversationID is the conversation that contains the message with the mention.
|
||||
ConversationID string
|
||||
// MessageID is the message that contains the mention.
|
||||
MessageID string
|
||||
// Kind describes the type of entity being mentioned (e.g. "conv", "thread", "doc").
|
||||
Kind string
|
||||
// TargetID is the UUID of the entity being mentioned.
|
||||
TargetID string
|
||||
// Raw is the raw mention token as it appeared in the message (e.g. "@conv:abc123").
|
||||
Raw string
|
||||
// Metadata is optional JSON-serialisable additional context.
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
// Add records a new mention.
|
||||
func (s *Store) Add(ctx context.Context, p AddParams) (sqlc.AgentMention, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return sqlc.AgentMention{}, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentMention{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
|
||||
msgIDStr := strings.TrimSpace(p.MessageID)
|
||||
if msgIDStr == "" {
|
||||
return sqlc.AgentMention{}, pcerrors.New(pcerrors.CodeInvalidArgument, "message_id is required")
|
||||
}
|
||||
msgID, err := ids.Parse(msgIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentMention{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse message_id %q", msgIDStr)
|
||||
}
|
||||
|
||||
kind := strings.TrimSpace(p.Kind)
|
||||
if kind == "" {
|
||||
return sqlc.AgentMention{}, pcerrors.New(pcerrors.CodeInvalidArgument, "kind is required")
|
||||
}
|
||||
|
||||
targetIDStr := strings.TrimSpace(p.TargetID)
|
||||
if targetIDStr == "" {
|
||||
return sqlc.AgentMention{}, pcerrors.New(pcerrors.CodeInvalidArgument, "target_id is required")
|
||||
}
|
||||
targetID, err := ids.Parse(targetIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentMention{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse target_id %q", targetIDStr)
|
||||
}
|
||||
|
||||
metaJSON, _ := json.Marshal(p.Metadata)
|
||||
|
||||
return s.q.AddAgentMention(ctx, sqlc.AddAgentMentionParams{
|
||||
ID: ids.New(),
|
||||
ConversationID: convID,
|
||||
MessageID: msgID,
|
||||
Kind: kind,
|
||||
TargetID: targetID,
|
||||
Raw: strings.TrimSpace(p.Raw),
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
}
|
||||
|
||||
// ─── ListByConversation ───────────────────────────────────────────────────────
|
||||
|
||||
// ListByConversationParams configures the ListByConversation operation.
|
||||
type ListByConversationParams struct {
|
||||
ConversationID string
|
||||
}
|
||||
|
||||
// ListByConversation returns all mentions recorded within a conversation.
|
||||
func (s *Store) ListByConversation(ctx context.Context, p ListByConversationParams) ([]sqlc.AgentMention, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return nil, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return nil, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
return s.q.ListAgentMentionsByConversationID(ctx,
|
||||
sqlc.ListAgentMentionsByConversationIDParams{ConversationID: convID})
|
||||
}
|
||||
166
pkg/agent/threads/store.go
Normal file
166
pkg/agent/threads/store.go
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
// Package threads provides the business-logic operations layer for agent
|
||||
// sub-threads within a conversation: creation, listing, and message management.
|
||||
// Sub-threads allow parallel or branching dialogue tracks without forking the
|
||||
// parent conversation.
|
||||
package threads
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/ids"
|
||||
sqlc "github.com/sipeed/picoclaw/pkg/memory/sqlc"
|
||||
"github.com/sipeed/picoclaw/pkg/pcerrors"
|
||||
)
|
||||
|
||||
// Store wraps the SQLC queries for thread operations.
|
||||
// Construct with New.
|
||||
type Store struct {
|
||||
q *sqlc.Queries
|
||||
}
|
||||
|
||||
// New returns a Store backed by the provided SQLC queries handle.
|
||||
func New(q *sqlc.Queries) *Store {
|
||||
return &Store{q: q}
|
||||
}
|
||||
|
||||
// ─── Create ──────────────────────────────────────────────────────────────────
|
||||
|
||||
// CreateParams configures the Create operation.
|
||||
type CreateParams struct {
|
||||
ConversationID string
|
||||
Title *string
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
// Create creates a new thread within a conversation.
|
||||
func (s *Store) Create(ctx context.Context, p CreateParams) (sqlc.AgentThread, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return sqlc.AgentThread{}, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentThread{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
|
||||
metaJSON, _ := json.Marshal(p.Metadata)
|
||||
|
||||
return s.q.CreateAgentThread(ctx, sqlc.CreateAgentThreadParams{
|
||||
ID: ids.New(),
|
||||
ConversationID: convID,
|
||||
Title: p.Title,
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
}
|
||||
|
||||
// ─── List ────────────────────────────────────────────────────────────────────
|
||||
|
||||
// ListParams configures the List operation.
|
||||
type ListParams struct {
|
||||
ConversationID string
|
||||
}
|
||||
|
||||
// List returns all threads for the given conversation.
|
||||
func (s *Store) List(ctx context.Context, p ListParams) ([]sqlc.AgentThread, error) {
|
||||
convIDStr := strings.TrimSpace(p.ConversationID)
|
||||
if convIDStr == "" {
|
||||
return nil, pcerrors.New(pcerrors.CodeInvalidArgument, "conversation_id is required")
|
||||
}
|
||||
convID, err := ids.Parse(convIDStr)
|
||||
if err != nil {
|
||||
return nil, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse conversation_id %q", convIDStr)
|
||||
}
|
||||
return s.q.ListAgentThreadsByConversationID(ctx,
|
||||
sqlc.ListAgentThreadsByConversationIDParams{ConversationID: convID})
|
||||
}
|
||||
|
||||
// ─── AddMessage ──────────────────────────────────────────────────────────────
|
||||
|
||||
// AddMessageParams configures the AddMessage operation.
|
||||
type AddMessageParams struct {
|
||||
ThreadID string
|
||||
// Role defaults to "user" if empty.
|
||||
Role string
|
||||
Content string
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
// AddMessage appends a message to a thread.
|
||||
func (s *Store) AddMessage(ctx context.Context, p AddMessageParams) (sqlc.AgentThreadMessage, error) {
|
||||
threadIDStr := strings.TrimSpace(p.ThreadID)
|
||||
if threadIDStr == "" {
|
||||
return sqlc.AgentThreadMessage{}, pcerrors.New(pcerrors.CodeInvalidArgument, "thread_id is required")
|
||||
}
|
||||
threadID, err := ids.Parse(threadIDStr)
|
||||
if err != nil {
|
||||
return sqlc.AgentThreadMessage{}, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse thread_id %q", threadIDStr)
|
||||
}
|
||||
|
||||
role := strings.TrimSpace(p.Role)
|
||||
if role == "" {
|
||||
role = "user"
|
||||
}
|
||||
|
||||
content := strings.TrimSpace(p.Content)
|
||||
if content == "" {
|
||||
return sqlc.AgentThreadMessage{}, pcerrors.New(pcerrors.CodeInvalidArgument, "content is empty")
|
||||
}
|
||||
|
||||
metaJSON, _ := json.Marshal(p.Metadata)
|
||||
|
||||
return s.q.AddAgentThreadMessage(ctx, sqlc.AddAgentThreadMessageParams{
|
||||
ID: ids.New(),
|
||||
ThreadID: threadID,
|
||||
Role: role,
|
||||
Content: content,
|
||||
MetadataJson: metaJSON,
|
||||
})
|
||||
}
|
||||
|
||||
// ─── ListMessages ────────────────────────────────────────────────────────────
|
||||
|
||||
// ListMessagesParams configures the ListMessages operation.
|
||||
type ListMessagesParams struct {
|
||||
ThreadID string
|
||||
// Limit is clamped to [1, 500]; defaults to 50.
|
||||
Limit int
|
||||
}
|
||||
|
||||
// ListMessages returns messages for a thread in chronological order.
|
||||
// Internally fetches in descending order and reverses so the caller always
|
||||
// receives oldest-first.
|
||||
func (s *Store) ListMessages(ctx context.Context, p ListMessagesParams) ([]sqlc.AgentThreadMessage, error) {
|
||||
threadIDStr := strings.TrimSpace(p.ThreadID)
|
||||
if threadIDStr == "" {
|
||||
return nil, pcerrors.New(pcerrors.CodeInvalidArgument, "thread_id is required")
|
||||
}
|
||||
threadID, err := ids.Parse(threadIDStr)
|
||||
if err != nil {
|
||||
return nil, pcerrors.Wrapf(pcerrors.CodeInvalidArgument, err, "parse thread_id %q", threadIDStr)
|
||||
}
|
||||
|
||||
limit := int64(p.Limit)
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
if limit > 500 {
|
||||
limit = 500
|
||||
}
|
||||
|
||||
rows, err := s.q.ListAgentThreadMessagesByThreadIDDescLimit(ctx,
|
||||
sqlc.ListAgentThreadMessagesByThreadIDDescLimitParams{
|
||||
ThreadID: threadID,
|
||||
Limit: limit,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Reverse to chronological order.
|
||||
for i, j := 0, len(rows)-1; i < j; i, j = i+1, j-1 {
|
||||
rows[i], rows[j] = rows[j], rows[i]
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
|
@ -28,21 +28,46 @@ import (
|
|||
// concatenated node outputs for synthesis.
|
||||
type JoinerFunc func(ctx context.Context, systemPrompt, resultSummary string) (string, uint32, error)
|
||||
|
||||
// RLMExpandFunc processes oversized context through recursive decomposition.
|
||||
// It receives a session key, query, and context content; returns the
|
||||
// synthesised answer and token cost. This bridges the DAG executor to the
|
||||
// RLM engine without creating import cycles.
|
||||
type RLMExpandFunc func(ctx context.Context, sessionKey, query, contextContent string) (string, uint32, error)
|
||||
|
||||
// Executor runs a DAGPlan through the SecureBus with topological dispatch.
|
||||
type Executor struct {
|
||||
bus *securebus.Bus
|
||||
joiner JoinerFunc
|
||||
maxParallel int
|
||||
bus *securebus.Bus
|
||||
joiner JoinerFunc
|
||||
rlmExpand RLMExpandFunc
|
||||
rlmThresholdBytes int
|
||||
maxParallel int
|
||||
}
|
||||
|
||||
// ExecutorOption configures an Executor via the functional options pattern.
|
||||
type ExecutorOption func(*Executor)
|
||||
|
||||
// WithRLMExpander enables automatic RLM expansion for nodes whose output
|
||||
// exceeds threshold bytes.
|
||||
func WithRLMExpander(fn RLMExpandFunc, thresholdBytes int) ExecutorOption {
|
||||
return func(e *Executor) {
|
||||
e.rlmExpand = fn
|
||||
e.rlmThresholdBytes = thresholdBytes
|
||||
}
|
||||
}
|
||||
|
||||
// NewExecutor creates a DAG executor.
|
||||
// joiner is called after all nodes complete to synthesise the final answer.
|
||||
func NewExecutor(bus *securebus.Bus, joiner JoinerFunc) *Executor {
|
||||
return &Executor{
|
||||
bus: bus,
|
||||
joiner: joiner,
|
||||
maxParallel: runtime.GOMAXPROCS(0),
|
||||
func NewExecutor(bus *securebus.Bus, joiner JoinerFunc, opts ...ExecutorOption) *Executor {
|
||||
e := &Executor{
|
||||
bus: bus,
|
||||
joiner: joiner,
|
||||
rlmThresholdBytes: 8192,
|
||||
maxParallel: runtime.GOMAXPROCS(0),
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(e)
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
// ExecuteResult holds the output of a DAG execution.
|
||||
|
|
@ -141,7 +166,20 @@ func (e *Executor) Execute(ctx context.Context, sessionKey string, plan *itr.DAG
|
|||
if resp.IsError {
|
||||
ns.setResult(resp.Result, fmt.Errorf("node %s: %s", nodeID, resp.Result))
|
||||
} else {
|
||||
ns.setResult(resp.Result, nil)
|
||||
result := resp.Result
|
||||
// RLM expansion: if the result exceeds the threshold,
|
||||
// recursively decompose it via the RLM engine.
|
||||
if e.rlmExpand != nil && len(result) > e.rlmThresholdBytes {
|
||||
expanded, rlmTokens, rlmErr := e.rlmExpand(ctx, sessionKey,
|
||||
"Summarize and extract key information from this content", result)
|
||||
if rlmErr == nil {
|
||||
result = expanded
|
||||
tokensMu.Lock()
|
||||
totalTokens += rlmTokens
|
||||
tokensMu.Unlock()
|
||||
}
|
||||
}
|
||||
ns.setResult(result, nil)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
|
|
|||
244
pkg/itr/dag/executor_test.go
Normal file
244
pkg/itr/dag/executor_test.go
Normal file
|
|
@ -0,0 +1,244 @@
|
|||
package dag_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/itr"
|
||||
"github.com/sipeed/picoclaw/pkg/itr/dag"
|
||||
"github.com/sipeed/picoclaw/pkg/security/securebus"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func makeBus(t *testing.T, toolMap map[string]tools.Tool) *securebus.Bus {
|
||||
t.Helper()
|
||||
capLookup := func(name string) (tools.ToolCapabilities, bool) {
|
||||
tool, ok := toolMap[name]
|
||||
if !ok {
|
||||
return tools.ZeroCapabilities(), false
|
||||
}
|
||||
return tools.ExtractCapabilities(tool), true
|
||||
}
|
||||
executor := func(ctx context.Context, name string, args map[string]interface{}) *tools.ToolResult {
|
||||
tool, ok := toolMap[name]
|
||||
if !ok {
|
||||
return &tools.ToolResult{ForLLM: "tool not found: " + name, IsError: true}
|
||||
}
|
||||
return tool.Execute(ctx, args)
|
||||
}
|
||||
return securebus.New(securebus.DefaultBusConfig(), nil, capLookup, executor)
|
||||
}
|
||||
|
||||
type staticTool struct {
|
||||
name string
|
||||
result string
|
||||
}
|
||||
|
||||
func (s *staticTool) Name() string { return s.name }
|
||||
func (s *staticTool) Description() string { return "test" }
|
||||
func (s *staticTool) Parameters() map[string]interface{} {
|
||||
return map[string]interface{}{"type": "object", "properties": map[string]interface{}{}}
|
||||
}
|
||||
func (s *staticTool) Execute(_ context.Context, args map[string]interface{}) *tools.ToolResult {
|
||||
if input, ok := args["input"].(string); ok {
|
||||
return &tools.ToolResult{ForLLM: fmt.Sprintf("%s:%s", s.result, input)}
|
||||
}
|
||||
return &tools.ToolResult{ForLLM: s.result}
|
||||
}
|
||||
|
||||
func TestExecutor_LinearDependencyChain(t *testing.T) {
|
||||
toolMap := map[string]tools.Tool{
|
||||
"step1": &staticTool{name: "step1", result: "r1"},
|
||||
"step2": &staticTool{name: "step2", result: "r2"},
|
||||
}
|
||||
bus := makeBus(t, toolMap)
|
||||
defer bus.Close()
|
||||
|
||||
executor := dag.NewExecutor(bus, nil)
|
||||
|
||||
plan := &itr.DAGPlan{
|
||||
Nodes: []itr.DAGNode{
|
||||
{ID: "a", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "step1", ArgsJSON: "{}"}},
|
||||
{ID: "b", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "step2", ArgsJSON: `{"input": "#nodea"}`}, DependsOn: []string{"a"}},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := executor.Execute(context.Background(), "test-sess", plan)
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, result.NodeResults["a"], "r1")
|
||||
assert.Contains(t, result.NodeResults["b"], "r2")
|
||||
}
|
||||
|
||||
func TestExecutor_ParallelNodes(t *testing.T) {
|
||||
toolMap := map[string]tools.Tool{
|
||||
"alpha": &staticTool{name: "alpha", result: "a-result"},
|
||||
"beta": &staticTool{name: "beta", result: "b-result"},
|
||||
}
|
||||
bus := makeBus(t, toolMap)
|
||||
defer bus.Close()
|
||||
|
||||
executor := dag.NewExecutor(bus, nil)
|
||||
|
||||
plan := &itr.DAGPlan{
|
||||
Nodes: []itr.DAGNode{
|
||||
{ID: "n1", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "alpha", ArgsJSON: "{}"}},
|
||||
{ID: "n2", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "beta", ArgsJSON: "{}"}},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := executor.Execute(context.Background(), "test-sess", plan)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "a-result", result.NodeResults["n1"])
|
||||
assert.Equal(t, "b-result", result.NodeResults["n2"])
|
||||
}
|
||||
|
||||
func TestExecutor_CycleDetection(t *testing.T) {
|
||||
bus := makeBus(t, map[string]tools.Tool{})
|
||||
defer bus.Close()
|
||||
|
||||
executor := dag.NewExecutor(bus, nil)
|
||||
|
||||
plan := &itr.DAGPlan{
|
||||
Nodes: []itr.DAGNode{
|
||||
{ID: "x", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "t", ArgsJSON: "{}"}, DependsOn: []string{"y"}},
|
||||
{ID: "y", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "t", ArgsJSON: "{}"}, DependsOn: []string{"x"}},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := executor.Execute(context.Background(), "test-sess", plan)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "cycle")
|
||||
}
|
||||
|
||||
func TestExecutor_WithJoiner(t *testing.T) {
|
||||
toolMap := map[string]tools.Tool{
|
||||
"tool1": &staticTool{name: "tool1", result: "data-A"},
|
||||
"tool2": &staticTool{name: "tool2", result: "data-B"},
|
||||
}
|
||||
bus := makeBus(t, toolMap)
|
||||
defer bus.Close()
|
||||
|
||||
joiner := func(_ context.Context, _, userQuery string) (string, uint32, error) {
|
||||
return "synthesized: " + userQuery[:20], 50, nil
|
||||
}
|
||||
|
||||
executor := dag.NewExecutor(bus, joiner)
|
||||
|
||||
plan := &itr.DAGPlan{
|
||||
Nodes: []itr.DAGNode{
|
||||
{ID: "n1", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "tool1", ArgsJSON: "{}"}},
|
||||
{ID: "n2", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "tool2", ArgsJSON: "{}"}},
|
||||
},
|
||||
JoinerQuery: "Combine the results into a summary",
|
||||
}
|
||||
|
||||
result, err := executor.Execute(context.Background(), "test-sess", plan)
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, result.FinalAnswer, "synthesized:")
|
||||
assert.Equal(t, uint32(50), result.TotalTokens)
|
||||
}
|
||||
|
||||
func TestExecutor_EmptyPlan(t *testing.T) {
|
||||
bus := makeBus(t, map[string]tools.Tool{})
|
||||
defer bus.Close()
|
||||
|
||||
executor := dag.NewExecutor(bus, nil)
|
||||
|
||||
result, err := executor.Execute(context.Background(), "test-sess", &itr.DAGPlan{})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, result.NodeResults)
|
||||
}
|
||||
|
||||
func TestResolver_NodeRefSubstitution(t *testing.T) {
|
||||
argsJSON := `{"query": "search for #nodeprev results"}`
|
||||
toolMap := map[string]tools.Tool{
|
||||
"search": &staticTool{name: "search", result: "found"},
|
||||
"prev": &staticTool{name: "prev", result: "previous-output"},
|
||||
}
|
||||
bus := makeBus(t, toolMap)
|
||||
defer bus.Close()
|
||||
|
||||
executor := dag.NewExecutor(bus, nil)
|
||||
|
||||
plan := &itr.DAGPlan{
|
||||
Nodes: []itr.DAGNode{
|
||||
{ID: "prev", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "prev", ArgsJSON: "{}"}},
|
||||
{ID: "search", Type: itr.CmdToolExec, Payload: itr.ToolExec{ToolName: "search", ArgsJSON: argsJSON}, DependsOn: []string{"prev"}},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := executor.Execute(context.Background(), "test-sess", plan)
|
||||
require.NoError(t, err)
|
||||
_ = result
|
||||
}
|
||||
|
||||
func TestRouter_SimpleQuerySelectsReAct(t *testing.T) {
|
||||
cfg := dag.DefaultRouterConfig()
|
||||
mode := dag.Route(dag.ModeAuto, "What is the weather?", cfg)
|
||||
assert.Equal(t, dag.ModeReAct, mode)
|
||||
}
|
||||
|
||||
func TestRouter_ComplexQuerySelectsDAG(t *testing.T) {
|
||||
cfg := dag.DefaultRouterConfig()
|
||||
mode := dag.Route(dag.ModeAuto, "Search for the latest news about AI, read the top 3 articles, and compare their viewpoints to create a summary report with aggregate statistics", cfg)
|
||||
assert.Equal(t, dag.ModeDAG, mode)
|
||||
}
|
||||
|
||||
func TestRouter_ExplicitModeOverridesAuto(t *testing.T) {
|
||||
cfg := dag.DefaultRouterConfig()
|
||||
mode := dag.Route(dag.ModeReAct, "Do many complex parallel things simultaneously", cfg)
|
||||
assert.Equal(t, dag.ModeReAct, mode)
|
||||
}
|
||||
|
||||
func TestPlanner_ValidatePlan(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
plan string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid simple plan",
|
||||
plan: `{"nodes":[{"id":"a","type":"tool_exec","payload":{"tool_name":"read_file","args_json":"{}"}}],"joiner_query":"summarize"}`,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "empty nodes",
|
||||
plan: `{"nodes":[],"joiner_query":"summarize"}`,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "duplicate IDs",
|
||||
plan: `{"nodes":[{"id":"a","type":"tool_exec","payload":{}},{"id":"a","type":"tool_exec","payload":{}}]}`,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "unknown dependency",
|
||||
plan: `{"nodes":[{"id":"a","type":"tool_exec","payload":{},"depends_on":["nonexistent"]}]}`,
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var plan itr.DAGPlan
|
||||
err := json.Unmarshal([]byte(tt.plan), &plan)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Use planner with a mock that returns the pre-built plan JSON
|
||||
mockModel := func(_ context.Context, _, _ string) (string, uint32, error) {
|
||||
return tt.plan, 10, nil
|
||||
}
|
||||
planner := dag.NewPlanner(mockModel, nil, dag.DefaultPlannerConfig())
|
||||
_, _, planErr := planner.Plan(context.Background(), "test query", nil)
|
||||
if tt.wantErr {
|
||||
assert.Error(t, planErr)
|
||||
} else {
|
||||
assert.NoError(t, planErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -64,6 +64,38 @@ func (q *Queries) AddAgentMessage(ctx context.Context, arg AddAgentMessageParams
|
|||
return i, err
|
||||
}
|
||||
|
||||
const GetAgentMessageByID = `-- name: GetAgentMessageByID :one
|
||||
SELECT id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
FROM agent_messages
|
||||
WHERE id = ?
|
||||
LIMIT 1
|
||||
`
|
||||
|
||||
type GetAgentMessageByIDParams struct {
|
||||
ID ids.UUID `db:"id" json:"id"`
|
||||
}
|
||||
|
||||
// GetAgentMessageByID
|
||||
//
|
||||
// SELECT id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
// FROM agent_messages
|
||||
// WHERE id = ?
|
||||
// LIMIT 1
|
||||
func (q *Queries) GetAgentMessageByID(ctx context.Context, arg GetAgentMessageByIDParams) (AgentMessage, error) {
|
||||
row := q.db.QueryRowContext(ctx, GetAgentMessageByID, arg.ID)
|
||||
var i AgentMessage
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.ConversationID,
|
||||
&i.Role,
|
||||
&i.Content,
|
||||
&i.MetadataJson,
|
||||
&i.CreatedAt,
|
||||
&i.UpdatedAt,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const ListAgentMessagesByConversationID = `-- name: ListAgentMessagesByConversationID :many
|
||||
SELECT id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
FROM agent_messages
|
||||
|
|
@ -162,3 +194,36 @@ func (q *Queries) ListAgentMessagesByConversationIDLimit(ctx context.Context, ar
|
|||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const UpdateAgentMessageContent = `-- name: UpdateAgentMessageContent :one
|
||||
UPDATE agent_messages
|
||||
SET content = ?
|
||||
WHERE id = ?
|
||||
RETURNING id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
`
|
||||
|
||||
type UpdateAgentMessageContentParams struct {
|
||||
Content string `db:"content" json:"content"`
|
||||
ID ids.UUID `db:"id" json:"id"`
|
||||
}
|
||||
|
||||
// UpdateAgentMessageContent
|
||||
//
|
||||
// UPDATE agent_messages
|
||||
// SET content = ?
|
||||
// WHERE id = ?
|
||||
// RETURNING id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
func (q *Queries) UpdateAgentMessageContent(ctx context.Context, arg UpdateAgentMessageContentParams) (AgentMessage, error) {
|
||||
row := q.db.QueryRowContext(ctx, UpdateAgentMessageContent, arg.Content, arg.ID)
|
||||
var i AgentMessage
|
||||
err := row.Scan(
|
||||
&i.ID,
|
||||
&i.ConversationID,
|
||||
&i.Role,
|
||||
&i.Content,
|
||||
&i.MetadataJson,
|
||||
&i.CreatedAt,
|
||||
&i.UpdatedAt,
|
||||
)
|
||||
return i, err
|
||||
}
|
||||
|
|
|
|||
|
|
@ -303,6 +303,13 @@ type Querier interface {
|
|||
// WHERE child_conversation_id = ?
|
||||
// LIMIT 1
|
||||
GetAgentConversationForkByChildConversationID(ctx context.Context, arg GetAgentConversationForkByChildConversationIDParams) (AgentConversationFork, error)
|
||||
//GetAgentMessageByID
|
||||
//
|
||||
// SELECT id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
// FROM agent_messages
|
||||
// WHERE id = ?
|
||||
// LIMIT 1
|
||||
GetAgentMessageByID(ctx context.Context, arg GetAgentMessageByIDParams) (AgentMessage, error)
|
||||
//GetAgentRunStateByID
|
||||
//
|
||||
// SELECT id, run_id, step_index, state, snapshot_json, created_at, updated_at
|
||||
|
|
@ -982,6 +989,13 @@ type Querier interface {
|
|||
// WHERE id = ?
|
||||
// RETURNING id, title, created_at, updated_at
|
||||
UpdateAgentConversationTitle(ctx context.Context, arg UpdateAgentConversationTitleParams) (AgentConversation, error)
|
||||
//UpdateAgentMessageContent
|
||||
//
|
||||
// UPDATE agent_messages
|
||||
// SET content = ?
|
||||
// WHERE id = ?
|
||||
// RETURNING id, conversation_id, role, content, metadata_json, created_at, updated_at
|
||||
UpdateAgentMessageContent(ctx context.Context, arg UpdateAgentMessageContentParams) (AgentMessage, error)
|
||||
//UpdateAgentRunStatus
|
||||
//
|
||||
// UPDATE agent_runs
|
||||
|
|
|
|||
|
|
@ -19,3 +19,13 @@ FROM agent_messages
|
|||
WHERE conversation_id = ?
|
||||
ORDER BY created_at ASC
|
||||
LIMIT ?;
|
||||
-- name: GetAgentMessageByID :one
|
||||
SELECT *
|
||||
FROM agent_messages
|
||||
WHERE id = ?
|
||||
LIMIT 1;
|
||||
-- name: UpdateAgentMessageContent :one
|
||||
UPDATE agent_messages
|
||||
SET content = ?
|
||||
WHERE id = ?
|
||||
RETURNING *;
|
||||
|
|
@ -20,6 +20,10 @@ type ToolExecutor func(ctx context.Context, name string, args map[string]interfa
|
|||
// Wraps tools.Registry.Get + tools.ExtractCapabilities.
|
||||
type CapabilitiesLookup func(toolName string) (tools.ToolCapabilities, bool)
|
||||
|
||||
// ToolSearchFunc searches the tool registry and returns matching tool info
|
||||
// as a JSON string. If nil, ToolSearch commands return an error.
|
||||
type ToolSearchFunc func(query string, maxResults int) string
|
||||
|
||||
// BusConfig configures the SecureBus.
|
||||
type BusConfig struct {
|
||||
Policy PolicyConfig
|
||||
|
|
@ -48,15 +52,16 @@ func DefaultBusConfig() BusConfig {
|
|||
// 6. Write audit log entry
|
||||
// 7. Return ToolResponse to caller
|
||||
type Bus struct {
|
||||
cfg BusConfig
|
||||
policy *PolicyEngine
|
||||
secrets *security.SecretStore // nil = no secret injection
|
||||
redactor *security.Redactor
|
||||
audit *AuditLog
|
||||
transport *ChannelTransport
|
||||
capLookup CapabilitiesLookup
|
||||
executor ToolExecutor
|
||||
done chan struct{}
|
||||
cfg BusConfig
|
||||
policy *PolicyEngine
|
||||
secrets *security.SecretStore // nil = no secret injection
|
||||
redactor *security.Redactor
|
||||
audit *AuditLog
|
||||
transport *ChannelTransport
|
||||
capLookup CapabilitiesLookup
|
||||
executor ToolExecutor
|
||||
toolSearch ToolSearchFunc // nil = no tool search support
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
// New creates a Bus and starts background worker goroutines.
|
||||
|
|
@ -142,9 +147,17 @@ func (b *Bus) dispatch(ctx context.Context, req itr.ToolRequest) itr.ToolRespons
|
|||
At: start,
|
||||
}
|
||||
|
||||
// Only ToolExec requests require capability/secret/leak checks.
|
||||
// RLM operations (Peek, Grep, etc.) are structural and access no tools.
|
||||
if req.Type != itr.CmdToolExec {
|
||||
// Only ToolExec requests require full capability/secret/leak checks.
|
||||
// ToolSearch, DAGPlan, and RLM operations are handled separately.
|
||||
switch req.Type {
|
||||
case itr.CmdToolExec:
|
||||
// Falls through to the capability/secret/leak pipeline below.
|
||||
case itr.CmdToolSearch:
|
||||
resp := b.handleToolSearch(ctx, req)
|
||||
event.DurationMS = time.Since(start).Milliseconds()
|
||||
_ = b.audit.Append(event)
|
||||
return resp
|
||||
default:
|
||||
resp := b.handleRLMCommand(ctx, req)
|
||||
event.DurationMS = time.Since(start).Milliseconds()
|
||||
_ = b.audit.Append(event)
|
||||
|
|
@ -261,13 +274,34 @@ func injectArg(args map[string]interface{}, injectAs, value string) {
|
|||
// recorded in the capability manifest so auditing can trace what was accessed.
|
||||
}
|
||||
|
||||
// SetToolSearch configures the tool search callback. Call this after
|
||||
// constructing the Bus if tool search is needed.
|
||||
func (b *Bus) SetToolSearch(fn ToolSearchFunc) {
|
||||
b.toolSearch = fn
|
||||
}
|
||||
|
||||
// handleToolSearch processes CmdToolSearch requests by delegating to the
|
||||
// configured ToolSearchFunc.
|
||||
func (b *Bus) handleToolSearch(_ context.Context, req itr.ToolRequest) itr.ToolResponse {
|
||||
ts, ok := req.Payload.(itr.ToolSearch)
|
||||
if !ok {
|
||||
return itr.NewErrorResponse(req.ID, "internal: payload is not ToolSearch")
|
||||
}
|
||||
if b.toolSearch == nil {
|
||||
return itr.NewErrorResponse(req.ID, "tool search not configured")
|
||||
}
|
||||
maxResults := int(ts.MaxResults)
|
||||
if maxResults <= 0 {
|
||||
maxResults = 10
|
||||
}
|
||||
result := b.toolSearch(ts.Query, maxResults)
|
||||
return itr.NewSuccessResponse(req.ID, result, 0)
|
||||
}
|
||||
|
||||
// handleRLMCommand processes structural RLM decomposition commands.
|
||||
// These commands don't invoke tool code — they operate on the context rope
|
||||
// managed by the RLMEngine (which calls the SecureBus, not the other way around).
|
||||
func (b *Bus) handleRLMCommand(_ context.Context, req itr.ToolRequest) itr.ToolResponse {
|
||||
// RLM commands are executed by the RLMEngine; if they reach the SecureBus
|
||||
// directly it means the engine called Bus.Execute with an RLM payload.
|
||||
// Return a stub response — the RLMEngine interprets this.
|
||||
switch req.Type {
|
||||
case itr.CmdFinal:
|
||||
if f, ok := req.Payload.(itr.Final); ok {
|
||||
|
|
|
|||
231
pkg/security/securebus/socket_transport.go
Normal file
231
pkg/security/securebus/socket_transport.go
Normal file
|
|
@ -0,0 +1,231 @@
|
|||
package securebus
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/itr"
|
||||
)
|
||||
|
||||
const (
|
||||
maxFrameSize = 16 * 1024 * 1024 // 16 MiB sanity limit
|
||||
socketFilePerms = 0600
|
||||
)
|
||||
|
||||
// SocketTransport implements Transport over a Unix domain socket using
|
||||
// length-prefixed JSON frames (4-byte big-endian length + JSON payload).
|
||||
// The server side listens for connections and dispatches requests to the
|
||||
// SecureBus; the client side connects and performs request/response exchanges.
|
||||
type SocketTransport struct {
|
||||
mu sync.Mutex
|
||||
path string
|
||||
conn net.Conn
|
||||
listener net.Listener
|
||||
closed chan struct{}
|
||||
isServer bool
|
||||
|
||||
connsMu sync.Mutex
|
||||
conns []net.Conn
|
||||
}
|
||||
|
||||
// NewSocketTransportClient connects to the daemon's Unix socket at path.
|
||||
func NewSocketTransportClient(path string) (*SocketTransport, error) {
|
||||
conn, err := net.DialTimeout("unix", path, 5*time.Second)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("connect to daemon at %s: %w", path, err)
|
||||
}
|
||||
return &SocketTransport{
|
||||
path: path,
|
||||
conn: conn,
|
||||
closed: make(chan struct{}),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// NewSocketTransportServer creates a listening Unix socket at path.
|
||||
// Call Serve() to start accepting connections.
|
||||
func NewSocketTransportServer(path string) (*SocketTransport, error) {
|
||||
_ = os.Remove(path)
|
||||
|
||||
listener, err := net.Listen("unix", path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("listen on %s: %w", path, err)
|
||||
}
|
||||
|
||||
if err := os.Chmod(path, socketFilePerms); err != nil {
|
||||
listener.Close()
|
||||
return nil, fmt.Errorf("chmod socket: %w", err)
|
||||
}
|
||||
|
||||
return &SocketTransport{
|
||||
path: path,
|
||||
listener: listener,
|
||||
closed: make(chan struct{}),
|
||||
isServer: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Send submits a request over the socket and blocks until the response arrives.
|
||||
// Client-side only.
|
||||
func (st *SocketTransport) Send(ctx context.Context, req itr.ToolRequest) (itr.ToolResponse, error) {
|
||||
if st.isServer {
|
||||
return itr.ToolResponse{}, fmt.Errorf("Send called on server transport; use Serve instead")
|
||||
}
|
||||
|
||||
st.mu.Lock()
|
||||
defer st.mu.Unlock()
|
||||
|
||||
select {
|
||||
case <-st.closed:
|
||||
return itr.ToolResponse{}, fmt.Errorf("transport closed")
|
||||
default:
|
||||
}
|
||||
|
||||
if err := writeFrame(st.conn, req); err != nil {
|
||||
return itr.ToolResponse{}, fmt.Errorf("write request: %w", err)
|
||||
}
|
||||
|
||||
var resp itr.ToolResponse
|
||||
if err := readFrame(st.conn, &resp); err != nil {
|
||||
return itr.ToolResponse{}, fmt.Errorf("read response: %w", err)
|
||||
}
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// Serve accepts connections and dispatches requests to handler. Blocks until
|
||||
// Close is called or the listener errors. Server-side only.
|
||||
func (st *SocketTransport) Serve(handler func(ctx context.Context, req itr.ToolRequest) itr.ToolResponse) error {
|
||||
if !st.isServer {
|
||||
return fmt.Errorf("Serve called on client transport")
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
defer wg.Wait()
|
||||
|
||||
for {
|
||||
conn, err := st.listener.Accept()
|
||||
if err != nil {
|
||||
select {
|
||||
case <-st.closed:
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("accept: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
st.connsMu.Lock()
|
||||
st.conns = append(st.conns, conn)
|
||||
st.connsMu.Unlock()
|
||||
|
||||
wg.Add(1)
|
||||
go func(c net.Conn) {
|
||||
defer wg.Done()
|
||||
defer c.Close()
|
||||
st.handleConnection(c, handler)
|
||||
}(conn)
|
||||
}
|
||||
}
|
||||
|
||||
func (st *SocketTransport) handleConnection(conn net.Conn, handler func(ctx context.Context, req itr.ToolRequest) itr.ToolResponse) {
|
||||
for {
|
||||
select {
|
||||
case <-st.closed:
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
var req itr.ToolRequest
|
||||
if err := readFrame(conn, &req); err != nil {
|
||||
if err == io.EOF {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
resp := handler(context.Background(), req)
|
||||
|
||||
if err := writeFrame(conn, resp); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Path returns the socket file path.
|
||||
func (st *SocketTransport) Path() string {
|
||||
return st.path
|
||||
}
|
||||
|
||||
// Close shuts down the transport.
|
||||
func (st *SocketTransport) Close() error {
|
||||
select {
|
||||
case <-st.closed:
|
||||
return nil
|
||||
default:
|
||||
close(st.closed)
|
||||
}
|
||||
|
||||
if st.listener != nil {
|
||||
st.listener.Close()
|
||||
_ = os.Remove(st.path)
|
||||
}
|
||||
|
||||
st.connsMu.Lock()
|
||||
for _, c := range st.conns {
|
||||
c.Close()
|
||||
}
|
||||
st.conns = nil
|
||||
st.connsMu.Unlock()
|
||||
|
||||
if st.conn != nil {
|
||||
st.conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeFrame writes a length-prefixed JSON frame to w.
|
||||
func writeFrame(w io.Writer, v interface{}) error {
|
||||
data, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal: %w", err)
|
||||
}
|
||||
|
||||
if len(data) > maxFrameSize {
|
||||
return fmt.Errorf("frame too large: %d > %d", len(data), maxFrameSize)
|
||||
}
|
||||
|
||||
var header [4]byte
|
||||
binary.BigEndian.PutUint32(header[:], uint32(len(data)))
|
||||
|
||||
if _, err := w.Write(header[:]); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = w.Write(data)
|
||||
return err
|
||||
}
|
||||
|
||||
// readFrame reads a length-prefixed JSON frame from r into v.
|
||||
func readFrame(r io.Reader, v interface{}) error {
|
||||
var header [4]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
size := binary.BigEndian.Uint32(header[:])
|
||||
if size > maxFrameSize {
|
||||
return fmt.Errorf("frame too large: %d > %d", size, maxFrameSize)
|
||||
}
|
||||
|
||||
buf := make([]byte, size)
|
||||
if _, err := io.ReadFull(r, buf); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return json.Unmarshal(buf, v)
|
||||
}
|
||||
125
pkg/security/securebus/socket_transport_test.go
Normal file
125
pkg/security/securebus/socket_transport_test.go
Normal file
|
|
@ -0,0 +1,125 @@
|
|||
package securebus
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/itr"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSocketTransportRoundTrip(t *testing.T) {
|
||||
sockPath := filepath.Join(t.TempDir(), "test.sock")
|
||||
|
||||
server, err := NewSocketTransportServer(sockPath)
|
||||
require.NoError(t, err)
|
||||
defer server.Close()
|
||||
|
||||
handler := func(ctx context.Context, req itr.ToolRequest) itr.ToolResponse {
|
||||
return itr.ToolResponse{
|
||||
ID: req.ID,
|
||||
Result: `{"echo":"` + req.ID + `"}`,
|
||||
}
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_ = server.Serve(handler)
|
||||
}()
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
client, err := NewSocketTransportClient(sockPath)
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
req := itr.ToolRequest{
|
||||
ID: "req-001",
|
||||
Type: itr.CmdToolExec,
|
||||
Payload: itr.ToolExec{ToolName: "echo", ArgsJSON: `{}`},
|
||||
Timestamp: time.Now().UnixNano(),
|
||||
}
|
||||
|
||||
resp, err := client.Send(context.Background(), req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "req-001", resp.ID)
|
||||
assert.Contains(t, resp.Result, "req-001")
|
||||
|
||||
server.Close()
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestSocketTransportMultipleRequests(t *testing.T) {
|
||||
sockPath := filepath.Join(t.TempDir(), "multi.sock")
|
||||
|
||||
server, err := NewSocketTransportServer(sockPath)
|
||||
require.NoError(t, err)
|
||||
defer server.Close()
|
||||
|
||||
handler := func(ctx context.Context, req itr.ToolRequest) itr.ToolResponse {
|
||||
return itr.ToolResponse{ID: req.ID, Result: req.ID + "-done"}
|
||||
}
|
||||
|
||||
go func() { _ = server.Serve(handler) }()
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
client, err := NewSocketTransportClient(sockPath)
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
req := itr.ToolRequest{
|
||||
ID: "req-" + string(rune('A'+i)),
|
||||
Type: itr.CmdToolExec,
|
||||
Timestamp: time.Now().UnixNano(),
|
||||
}
|
||||
resp, err := client.Send(context.Background(), req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, req.ID, resp.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketTransportCleanup(t *testing.T) {
|
||||
sockPath := filepath.Join(t.TempDir(), "cleanup.sock")
|
||||
|
||||
server, err := NewSocketTransportServer(sockPath)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, statErr := os.Stat(sockPath)
|
||||
assert.NoError(t, statErr, "socket file should exist")
|
||||
|
||||
server.Close()
|
||||
|
||||
_, statErr = os.Stat(sockPath)
|
||||
assert.True(t, os.IsNotExist(statErr), "socket file should be removed after Close")
|
||||
}
|
||||
|
||||
func TestSocketTransportClientSendOnClosed(t *testing.T) {
|
||||
sockPath := filepath.Join(t.TempDir(), "closed.sock")
|
||||
|
||||
server, err := NewSocketTransportServer(sockPath)
|
||||
require.NoError(t, err)
|
||||
go func() {
|
||||
_ = server.Serve(func(ctx context.Context, req itr.ToolRequest) itr.ToolResponse {
|
||||
return itr.ToolResponse{ID: req.ID}
|
||||
})
|
||||
}()
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
client, err := NewSocketTransportClient(sockPath)
|
||||
require.NoError(t, err)
|
||||
|
||||
client.Close()
|
||||
|
||||
_, err = client.Send(context.Background(), itr.ToolRequest{ID: "fail"})
|
||||
assert.Error(t, err)
|
||||
|
||||
server.Close()
|
||||
}
|
||||
307
pkg/security/zkp.go
Normal file
307
pkg/security/zkp.go
Normal file
|
|
@ -0,0 +1,307 @@
|
|||
// Package security provides the Schnorr ZKP session handshake for daemon
|
||||
// authentication. The protocol proves knowledge of a shared secret (derived
|
||||
// from the master key) without revealing it, in a single round-trip (~200 bytes).
|
||||
//
|
||||
// Protocol (Schnorr identification on P-256):
|
||||
//
|
||||
// Prover Verifier
|
||||
// ────── ────────
|
||||
// k ← rand; R = k·G → (commitment)
|
||||
// ← c (32-byte challenge)
|
||||
// s = k − c·x mod n → (response)
|
||||
// s·G + c·Y == R ?
|
||||
//
|
||||
// On success the verifier issues a session token (random 32 bytes, TTL 1h).
|
||||
package security
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
var curve = elliptic.P256()
|
||||
|
||||
// SchnorrKeypair derives a Schnorr keypair from a 32-byte master key.
|
||||
// x = SHA-256(masterKey || "schnorr-zkp") mod n
|
||||
// Y = x·G
|
||||
func SchnorrKeypair(masterKey []byte) (*big.Int, *ecdsa.PublicKey, error) {
|
||||
if len(masterKey) != 32 {
|
||||
return nil, nil, fmt.Errorf("master key must be 32 bytes, got %d", len(masterKey))
|
||||
}
|
||||
|
||||
h := sha256.New()
|
||||
h.Write(masterKey)
|
||||
h.Write([]byte("schnorr-zkp"))
|
||||
xBytes := h.Sum(nil)
|
||||
|
||||
x := new(big.Int).SetBytes(xBytes)
|
||||
x.Mod(x, curve.Params().N)
|
||||
|
||||
if x.Sign() == 0 {
|
||||
return nil, nil, fmt.Errorf("degenerate key (zero scalar)")
|
||||
}
|
||||
|
||||
px, py := curve.ScalarBaseMult(x.Bytes())
|
||||
pub := &ecdsa.PublicKey{
|
||||
Curve: curve,
|
||||
X: px,
|
||||
Y: py,
|
||||
}
|
||||
return x, pub, nil
|
||||
}
|
||||
|
||||
// SchnorrCommitment is the prover's initial message.
|
||||
type SchnorrCommitment struct {
|
||||
RX, RY []byte // compressed point R = k·G
|
||||
k *big.Int
|
||||
}
|
||||
|
||||
// ProverCommit generates a random nonce and returns the commitment.
|
||||
func ProverCommit() (*SchnorrCommitment, error) {
|
||||
k, err := rand.Int(rand.Reader, curve.Params().N)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generate nonce: %w", err)
|
||||
}
|
||||
|
||||
rx, ry := curve.ScalarBaseMult(k.Bytes())
|
||||
return &SchnorrCommitment{
|
||||
RX: rx.Bytes(),
|
||||
RY: ry.Bytes(),
|
||||
k: k,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ProverRespond computes the response s = k - c·x mod n.
|
||||
func ProverRespond(commit *SchnorrCommitment, challenge []byte, secretKey *big.Int) ([]byte, error) {
|
||||
n := curve.Params().N
|
||||
c := new(big.Int).SetBytes(challenge)
|
||||
c.Mod(c, n)
|
||||
|
||||
cx := new(big.Int).Mul(c, secretKey)
|
||||
cx.Mod(cx, n)
|
||||
|
||||
s := new(big.Int).Sub(commit.k, cx)
|
||||
s.Mod(s, n)
|
||||
|
||||
sBytes := make([]byte, 32)
|
||||
sBuf := s.Bytes()
|
||||
copy(sBytes[32-len(sBuf):], sBuf)
|
||||
return sBytes, nil
|
||||
}
|
||||
|
||||
// VerifierChallenge generates a random 32-byte challenge.
|
||||
func VerifierChallenge() ([]byte, error) {
|
||||
c := make([]byte, 32)
|
||||
if _, err := rand.Read(c); err != nil {
|
||||
return nil, fmt.Errorf("generate challenge: %w", err)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// VerifierCheck verifies the Schnorr proof: s·G + c·Y == R.
|
||||
func VerifierCheck(pubKey *ecdsa.PublicKey, rx, ry, challenge, response []byte) bool {
|
||||
n := curve.Params().N
|
||||
|
||||
rX := new(big.Int).SetBytes(rx)
|
||||
rY := new(big.Int).SetBytes(ry)
|
||||
|
||||
if rX.Sign() == 0 && rY.Sign() == 0 {
|
||||
return false
|
||||
}
|
||||
if !curve.IsOnCurve(rX, rY) {
|
||||
return false
|
||||
}
|
||||
|
||||
s := new(big.Int).SetBytes(response)
|
||||
s.Mod(s, n)
|
||||
|
||||
c := new(big.Int).SetBytes(challenge)
|
||||
c.Mod(c, n)
|
||||
|
||||
// s·G
|
||||
sgx, sgy := curve.ScalarBaseMult(s.Bytes())
|
||||
|
||||
// c·Y
|
||||
cyx, cyy := curve.ScalarMult(pubKey.X, pubKey.Y, c.Bytes())
|
||||
|
||||
// s·G + c·Y
|
||||
checkX, checkY := curve.Add(sgx, sgy, cyx, cyy)
|
||||
|
||||
return checkX.Cmp(rX) == 0 && checkY.Cmp(rY) == 0
|
||||
}
|
||||
|
||||
// SessionToken is issued on successful ZKP handshake.
|
||||
type SessionToken struct {
|
||||
Token [32]byte
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// IsValid checks whether the token has not expired.
|
||||
func (st SessionToken) IsValid() bool {
|
||||
return time.Now().Before(st.ExpiresAt)
|
||||
}
|
||||
|
||||
// TokenHex returns the token as a hex string.
|
||||
func (st SessionToken) TokenHex() string {
|
||||
return fmt.Sprintf("%x", st.Token)
|
||||
}
|
||||
|
||||
// SessionManager tracks active session tokens for daemon auth.
|
||||
type ZKPSessionManager struct {
|
||||
mu sync.RWMutex
|
||||
pubKey *ecdsa.PublicKey
|
||||
sessions map[[32]byte]SessionToken
|
||||
ttl time.Duration
|
||||
}
|
||||
|
||||
// NewZKPSessionManager creates a session manager for the given public key.
|
||||
func NewZKPSessionManager(pubKey *ecdsa.PublicKey, ttl time.Duration) *ZKPSessionManager {
|
||||
if ttl == 0 {
|
||||
ttl = time.Hour
|
||||
}
|
||||
return &ZKPSessionManager{
|
||||
pubKey: pubKey,
|
||||
sessions: make(map[[32]byte]SessionToken),
|
||||
ttl: ttl,
|
||||
}
|
||||
}
|
||||
|
||||
// VerifyAndIssue performs the verifier side of the ZKP handshake.
|
||||
// On success, returns a new session token.
|
||||
func (sm *ZKPSessionManager) VerifyAndIssue(rx, ry, challenge, response []byte) (SessionToken, error) {
|
||||
if !VerifierCheck(sm.pubKey, rx, ry, challenge, response) {
|
||||
return SessionToken{}, errors.New("ZKP verification failed: invalid proof")
|
||||
}
|
||||
|
||||
var token [32]byte
|
||||
if _, err := rand.Read(token[:]); err != nil {
|
||||
return SessionToken{}, fmt.Errorf("generate session token: %w", err)
|
||||
}
|
||||
|
||||
st := SessionToken{
|
||||
Token: token,
|
||||
ExpiresAt: time.Now().Add(sm.ttl),
|
||||
}
|
||||
|
||||
sm.mu.Lock()
|
||||
sm.sessions[token] = st
|
||||
sm.mu.Unlock()
|
||||
|
||||
return st, nil
|
||||
}
|
||||
|
||||
// ValidateToken checks whether a token is known and not expired.
|
||||
func (sm *ZKPSessionManager) ValidateToken(token [32]byte) bool {
|
||||
sm.mu.RLock()
|
||||
st, ok := sm.sessions[token]
|
||||
sm.mu.RUnlock()
|
||||
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if !st.IsValid() {
|
||||
sm.mu.Lock()
|
||||
delete(sm.sessions, token)
|
||||
sm.mu.Unlock()
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// RevokeToken removes a session token.
|
||||
func (sm *ZKPSessionManager) RevokeToken(token [32]byte) {
|
||||
sm.mu.Lock()
|
||||
delete(sm.sessions, token)
|
||||
sm.mu.Unlock()
|
||||
}
|
||||
|
||||
// Cleanup removes all expired sessions.
|
||||
func (sm *ZKPSessionManager) Cleanup() int {
|
||||
sm.mu.Lock()
|
||||
defer sm.mu.Unlock()
|
||||
|
||||
expired := 0
|
||||
for k, st := range sm.sessions {
|
||||
if !st.IsValid() {
|
||||
delete(sm.sessions, k)
|
||||
expired++
|
||||
}
|
||||
}
|
||||
return expired
|
||||
}
|
||||
|
||||
// ActiveSessions returns the number of active sessions.
|
||||
func (sm *ZKPSessionManager) ActiveSessions() int {
|
||||
sm.mu.RLock()
|
||||
defer sm.mu.RUnlock()
|
||||
return len(sm.sessions)
|
||||
}
|
||||
|
||||
// HandshakePayload is the wire format for the ZKP handshake over the socket.
|
||||
// Total: 32 + 32 + 32 + 32 + 32 = 160 bytes (under 200-byte target).
|
||||
type HandshakePayload struct {
|
||||
RX [32]byte `json:"rx"`
|
||||
RY [32]byte `json:"ry"`
|
||||
Challenge [32]byte `json:"challenge"`
|
||||
Response [32]byte `json:"response"`
|
||||
Phase uint8 `json:"phase"` // 1=commit, 2=challenge, 3=response
|
||||
}
|
||||
|
||||
// MarshalBinary encodes the handshake payload as a compact binary frame.
|
||||
// Format: [1 phase][32 RX][32 RY][32 challenge][32 response] = 129 bytes
|
||||
func (hp HandshakePayload) MarshalBinary() []byte {
|
||||
buf := make([]byte, 129)
|
||||
buf[0] = hp.Phase
|
||||
copy(buf[1:33], hp.RX[:])
|
||||
copy(buf[33:65], hp.RY[:])
|
||||
copy(buf[65:97], hp.Challenge[:])
|
||||
copy(buf[97:129], hp.Response[:])
|
||||
return buf
|
||||
}
|
||||
|
||||
// UnmarshalBinaryHandshake decodes a compact binary handshake payload.
|
||||
func UnmarshalBinaryHandshake(data []byte) (HandshakePayload, error) {
|
||||
if len(data) < 129 {
|
||||
return HandshakePayload{}, fmt.Errorf("handshake payload too short: %d < 129", len(data))
|
||||
}
|
||||
var hp HandshakePayload
|
||||
hp.Phase = data[0]
|
||||
copy(hp.RX[:], data[1:33])
|
||||
copy(hp.RY[:], data[33:65])
|
||||
copy(hp.Challenge[:], data[65:97])
|
||||
copy(hp.Response[:], data[97:129])
|
||||
return hp, nil
|
||||
}
|
||||
|
||||
// HandshakeResult is the verifier's response after a successful handshake.
|
||||
type HandshakeResult struct {
|
||||
SessionToken [32]byte `json:"session_token"`
|
||||
ExpiresUnix int64 `json:"expires_unix"`
|
||||
}
|
||||
|
||||
// MarshalBinary encodes the result as [32 token][8 expires_unix] = 40 bytes.
|
||||
func (hr HandshakeResult) MarshalBinary() []byte {
|
||||
buf := make([]byte, 40)
|
||||
copy(buf[:32], hr.SessionToken[:])
|
||||
binary.BigEndian.PutUint64(buf[32:40], uint64(hr.ExpiresUnix))
|
||||
return buf
|
||||
}
|
||||
|
||||
// UnmarshalBinaryResult decodes a binary handshake result.
|
||||
func UnmarshalBinaryResult(data []byte) (HandshakeResult, error) {
|
||||
if len(data) < 40 {
|
||||
return HandshakeResult{}, fmt.Errorf("handshake result too short: %d < 40", len(data))
|
||||
}
|
||||
var hr HandshakeResult
|
||||
copy(hr.SessionToken[:], data[:32])
|
||||
hr.ExpiresUnix = int64(binary.BigEndian.Uint64(data[32:40]))
|
||||
return hr, nil
|
||||
}
|
||||
201
pkg/security/zkp_test.go
Normal file
201
pkg/security/zkp_test.go
Normal file
|
|
@ -0,0 +1,201 @@
|
|||
package security
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSchnorrKeypair(t *testing.T) {
|
||||
key := make([]byte, 32)
|
||||
for i := range key {
|
||||
key[i] = byte(i + 1)
|
||||
}
|
||||
|
||||
x, pub, err := SchnorrKeypair(key)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, x)
|
||||
assert.NotNil(t, pub)
|
||||
assert.True(t, curve.IsOnCurve(pub.X, pub.Y))
|
||||
}
|
||||
|
||||
func TestSchnorrKeypairRejectsBadLength(t *testing.T) {
|
||||
_, _, err := SchnorrKeypair([]byte("short"))
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestSchnorrFullHandshake(t *testing.T) {
|
||||
masterKey := make([]byte, 32)
|
||||
for i := range masterKey {
|
||||
masterKey[i] = byte(i + 42)
|
||||
}
|
||||
|
||||
x, pub, err := SchnorrKeypair(masterKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
commit, err := ProverCommit()
|
||||
require.NoError(t, err)
|
||||
|
||||
challenge, err := VerifierChallenge()
|
||||
require.NoError(t, err)
|
||||
|
||||
response, err := ProverRespond(commit, challenge, x)
|
||||
require.NoError(t, err)
|
||||
|
||||
valid := VerifierCheck(pub, commit.RX, commit.RY, challenge, response)
|
||||
assert.True(t, valid, "valid proof should verify")
|
||||
}
|
||||
|
||||
func TestSchnorrRejectsWrongKey(t *testing.T) {
|
||||
masterKey1 := make([]byte, 32)
|
||||
masterKey2 := make([]byte, 32)
|
||||
for i := range masterKey1 {
|
||||
masterKey1[i] = byte(i)
|
||||
masterKey2[i] = byte(i + 100)
|
||||
}
|
||||
|
||||
x1, _, err := SchnorrKeypair(masterKey1)
|
||||
require.NoError(t, err)
|
||||
_, pub2, err := SchnorrKeypair(masterKey2)
|
||||
require.NoError(t, err)
|
||||
|
||||
commit, err := ProverCommit()
|
||||
require.NoError(t, err)
|
||||
challenge, err := VerifierChallenge()
|
||||
require.NoError(t, err)
|
||||
response, err := ProverRespond(commit, challenge, x1)
|
||||
require.NoError(t, err)
|
||||
|
||||
valid := VerifierCheck(pub2, commit.RX, commit.RY, challenge, response)
|
||||
assert.False(t, valid, "proof with wrong key should fail")
|
||||
}
|
||||
|
||||
func TestZKPSessionManagerIssueAndValidate(t *testing.T) {
|
||||
masterKey := make([]byte, 32)
|
||||
for i := range masterKey {
|
||||
masterKey[i] = byte(i + 7)
|
||||
}
|
||||
|
||||
x, pub, err := SchnorrKeypair(masterKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
sm := NewZKPSessionManager(pub, time.Hour)
|
||||
|
||||
commit, _ := ProverCommit()
|
||||
challenge, _ := VerifierChallenge()
|
||||
response, _ := ProverRespond(commit, challenge, x)
|
||||
|
||||
st, err := sm.VerifyAndIssue(commit.RX, commit.RY, challenge, response)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, st.IsValid())
|
||||
assert.Equal(t, 1, sm.ActiveSessions())
|
||||
|
||||
assert.True(t, sm.ValidateToken(st.Token))
|
||||
}
|
||||
|
||||
func TestZKPSessionManagerRejectsInvalidProof(t *testing.T) {
|
||||
masterKey := make([]byte, 32)
|
||||
for i := range masterKey {
|
||||
masterKey[i] = byte(i)
|
||||
}
|
||||
_, pub, _ := SchnorrKeypair(masterKey)
|
||||
sm := NewZKPSessionManager(pub, time.Hour)
|
||||
|
||||
_, err := sm.VerifyAndIssue(make([]byte, 32), make([]byte, 32), make([]byte, 32), make([]byte, 32))
|
||||
assert.Error(t, err)
|
||||
assert.Equal(t, 0, sm.ActiveSessions())
|
||||
}
|
||||
|
||||
func TestZKPSessionManagerExpiry(t *testing.T) {
|
||||
masterKey := make([]byte, 32)
|
||||
for i := range masterKey {
|
||||
masterKey[i] = byte(i + 3)
|
||||
}
|
||||
x, pub, _ := SchnorrKeypair(masterKey)
|
||||
sm := NewZKPSessionManager(pub, 1*time.Millisecond)
|
||||
|
||||
commit, _ := ProverCommit()
|
||||
challenge, _ := VerifierChallenge()
|
||||
response, _ := ProverRespond(commit, challenge, x)
|
||||
|
||||
st, err := sm.VerifyAndIssue(commit.RX, commit.RY, challenge, response)
|
||||
require.NoError(t, err)
|
||||
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
assert.False(t, sm.ValidateToken(st.Token))
|
||||
}
|
||||
|
||||
func TestZKPSessionManagerRevoke(t *testing.T) {
|
||||
masterKey := make([]byte, 32)
|
||||
for i := range masterKey {
|
||||
masterKey[i] = byte(i + 5)
|
||||
}
|
||||
x, pub, _ := SchnorrKeypair(masterKey)
|
||||
sm := NewZKPSessionManager(pub, time.Hour)
|
||||
|
||||
commit, _ := ProverCommit()
|
||||
challenge, _ := VerifierChallenge()
|
||||
response, _ := ProverRespond(commit, challenge, x)
|
||||
|
||||
st, _ := sm.VerifyAndIssue(commit.RX, commit.RY, challenge, response)
|
||||
assert.True(t, sm.ValidateToken(st.Token))
|
||||
|
||||
sm.RevokeToken(st.Token)
|
||||
assert.False(t, sm.ValidateToken(st.Token))
|
||||
assert.Equal(t, 0, sm.ActiveSessions())
|
||||
}
|
||||
|
||||
func TestHandshakePayloadBinaryRoundTrip(t *testing.T) {
|
||||
hp := HandshakePayload{Phase: 3}
|
||||
for i := 0; i < 32; i++ {
|
||||
hp.RX[i] = byte(i)
|
||||
hp.RY[i] = byte(i + 32)
|
||||
hp.Challenge[i] = byte(i + 64)
|
||||
hp.Response[i] = byte(i + 96)
|
||||
}
|
||||
|
||||
data := hp.MarshalBinary()
|
||||
assert.Equal(t, 129, len(data))
|
||||
|
||||
decoded, err := UnmarshalBinaryHandshake(data)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, hp, decoded)
|
||||
}
|
||||
|
||||
func TestHandshakeResultBinaryRoundTrip(t *testing.T) {
|
||||
hr := HandshakeResult{ExpiresUnix: time.Now().Unix()}
|
||||
for i := 0; i < 32; i++ {
|
||||
hr.SessionToken[i] = byte(i + 200)
|
||||
}
|
||||
|
||||
data := hr.MarshalBinary()
|
||||
assert.Equal(t, 40, len(data))
|
||||
|
||||
decoded, err := UnmarshalBinaryResult(data)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, hr, decoded)
|
||||
}
|
||||
|
||||
func TestZKPSessionManagerCleanup(t *testing.T) {
|
||||
masterKey := make([]byte, 32)
|
||||
for i := range masterKey {
|
||||
masterKey[i] = byte(i + 11)
|
||||
}
|
||||
x, pub, _ := SchnorrKeypair(masterKey)
|
||||
sm := NewZKPSessionManager(pub, 1*time.Millisecond)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
commit, _ := ProverCommit()
|
||||
challenge, _ := VerifierChallenge()
|
||||
response, _ := ProverRespond(commit, challenge, x)
|
||||
_, _ = sm.VerifyAndIssue(commit.RX, commit.RY, challenge, response)
|
||||
}
|
||||
assert.Equal(t, 5, sm.ActiveSessions())
|
||||
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
cleaned := sm.Cleanup()
|
||||
assert.Equal(t, 5, cleaned)
|
||||
assert.Equal(t, 0, sm.ActiveSessions())
|
||||
}
|
||||
|
|
@ -16,6 +16,32 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/logger"
|
||||
)
|
||||
|
||||
// ToolLoopMode controls whether the agent uses the sequential ReAct loop,
|
||||
// the parallel DAG executor, or lets the router decide automatically.
|
||||
type ToolLoopMode int
|
||||
|
||||
const (
|
||||
ModeReAct ToolLoopMode = iota
|
||||
ModeDAG
|
||||
ModeAuto
|
||||
)
|
||||
|
||||
// DAGRunResult holds the output from a DAG execution.
|
||||
type DAGRunResult struct {
|
||||
Answer string
|
||||
Tokens uint32
|
||||
Iterations int
|
||||
}
|
||||
|
||||
// DAGRunFunc executes a query through the DAG planner/executor pipeline.
|
||||
// This function type breaks the import cycle between tools → dag → securebus → tools.
|
||||
// The concrete implementation is wired in the application entry point.
|
||||
type DAGRunFunc func(ctx context.Context, sessionKey, query string, availableTools []string) (*DAGRunResult, error)
|
||||
|
||||
// RouteFunc classifies a query and returns the preferred execution mode.
|
||||
// When nil, all queries use ModeReAct.
|
||||
type RouteFunc func(mode ToolLoopMode, query string) ToolLoopMode
|
||||
|
||||
// ToolLoopConfig configures the tool execution loop.
|
||||
type ToolLoopConfig struct {
|
||||
Model fantasy.LanguageModel
|
||||
|
|
@ -23,6 +49,17 @@ type ToolLoopConfig struct {
|
|||
Tools *ToolRegistry
|
||||
Bus *bus.MessageBus
|
||||
MaxIterations int
|
||||
|
||||
// DAGRunner executes queries through the DAG planner/executor pipeline.
|
||||
// When nil, all queries use the sequential ReAct loop.
|
||||
DAGRunner DAGRunFunc
|
||||
|
||||
// Router classifies queries into ModeReAct or ModeDAG. When nil,
|
||||
// ModeReAct is always used.
|
||||
Router RouteFunc
|
||||
|
||||
// LoopMode controls execution routing. Default: ModeAuto.
|
||||
LoopMode ToolLoopMode
|
||||
}
|
||||
|
||||
// ToolLoopResult contains the result of running the tool loop.
|
||||
|
|
@ -31,13 +68,32 @@ type ToolLoopResult struct {
|
|||
Iterations int
|
||||
}
|
||||
|
||||
// RunToolLoop executes the Fantasy agent loop with PicoClaw tools.
|
||||
// This is the core agent logic reused by both main agent and subagents.
|
||||
// RunToolLoop executes the agent loop with PicoClaw tools. It supports two
|
||||
// execution modes:
|
||||
// - ReAct (sequential): Fantasy's step-by-step tool calling loop
|
||||
// - DAG (parallel): LLMCompiler-style DAG planning and execution
|
||||
//
|
||||
// When LoopMode is ModeAuto, the router classifies the query to pick the
|
||||
// optimal mode. The SecureBus enforces capabilities in both modes.
|
||||
func RunToolLoop(ctx context.Context, config ToolLoopConfig, systemPrompt, userPrompt, channel, chatID string) (*ToolLoopResult, error) {
|
||||
// Build adapted tools
|
||||
mode := config.LoopMode
|
||||
if config.Router != nil {
|
||||
mode = config.Router(mode, userPrompt)
|
||||
} else if mode == ModeAuto {
|
||||
mode = ModeReAct
|
||||
}
|
||||
|
||||
if mode == ModeDAG && config.DAGRunner != nil {
|
||||
return runDAGLoop(ctx, config, userPrompt, channel)
|
||||
}
|
||||
|
||||
return runReActLoop(ctx, config, systemPrompt, userPrompt, channel, chatID)
|
||||
}
|
||||
|
||||
// runReActLoop is the original sequential Fantasy agent loop.
|
||||
func runReActLoop(ctx context.Context, config ToolLoopConfig, systemPrompt, userPrompt, channel, chatID string) (*ToolLoopResult, error) {
|
||||
adaptedTools := BuildAdaptedToolsFromRegistry(config.Tools, config.Bus, channel, chatID)
|
||||
|
||||
// Create Fantasy agent
|
||||
agentOpts := []fantasy.AgentOption{
|
||||
fantasy.WithTools(adaptedTools...),
|
||||
fantasy.WithStopConditions(fantasy.StepCountIs(config.MaxIterations)),
|
||||
|
|
@ -47,28 +103,25 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, systemPrompt, userP
|
|||
}
|
||||
agent := fantasy.NewAgent(config.Model, agentOpts...)
|
||||
|
||||
logger.DebugCF("toolloop", "Fantasy agent created for tool loop",
|
||||
logger.DebugCF("toolloop", "ReAct mode: Fantasy agent created",
|
||||
map[string]any{
|
||||
"tools_count": len(adaptedTools),
|
||||
"max_iterations": config.MaxIterations,
|
||||
})
|
||||
|
||||
// Run Fantasy agent
|
||||
result, err := agent.Generate(ctx, fantasy.AgentCall{
|
||||
Prompt: userPrompt,
|
||||
})
|
||||
if err != nil {
|
||||
logger.ErrorCF("toolloop", "Fantasy agent.Generate failed",
|
||||
map[string]any{
|
||||
"error": err.Error(),
|
||||
})
|
||||
map[string]any{"error": err.Error()})
|
||||
return nil, fmt.Errorf("agent Generate failed: %w", err)
|
||||
}
|
||||
|
||||
finalContent := result.Response.Content.Text()
|
||||
stepCount := len(result.Steps)
|
||||
|
||||
logger.InfoCF("toolloop", "Tool loop completed",
|
||||
logger.InfoCF("toolloop", "ReAct loop completed",
|
||||
map[string]any{
|
||||
"steps": stepCount,
|
||||
"content_chars": len(finalContent),
|
||||
|
|
@ -80,6 +133,33 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, systemPrompt, userP
|
|||
}, nil
|
||||
}
|
||||
|
||||
// runDAGLoop uses the LLMCompiler-style DAG executor with replanning.
|
||||
func runDAGLoop(ctx context.Context, config ToolLoopConfig, query, sessionKey string) (*ToolLoopResult, error) {
|
||||
logger.InfoCF("toolloop", "DAG mode: planning and executing",
|
||||
map[string]any{"query_len": len(query)})
|
||||
|
||||
availableTools := config.Tools.List()
|
||||
|
||||
result, err := config.DAGRunner(ctx, sessionKey, query, availableTools)
|
||||
if err != nil {
|
||||
logger.ErrorCF("toolloop", "DAG execution failed",
|
||||
map[string]any{"error": err.Error()})
|
||||
return nil, fmt.Errorf("DAG execution failed: %w", err)
|
||||
}
|
||||
|
||||
logger.InfoCF("toolloop", "DAG loop completed",
|
||||
map[string]any{
|
||||
"iterations": result.Iterations,
|
||||
"total_tokens": result.Tokens,
|
||||
"answer_chars": len(result.Answer),
|
||||
})
|
||||
|
||||
return &ToolLoopResult{
|
||||
Content: result.Answer,
|
||||
Iterations: result.Iterations,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// BuildAdaptedToolsFromRegistry wraps all tools in a ToolRegistry as Fantasy AgentTools.
|
||||
// This is a local wrapper that avoids circular imports by duplicating the adapter logic.
|
||||
func BuildAdaptedToolsFromRegistry(registry *ToolRegistry, msgBus *bus.MessageBus, channel, chatID string) []fantasy.AgentTool {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue