mirror of
https://github.com/google/ax.git
synced 2026-10-02 03:14:37 +08:00
Use cobra for CLI management
This commit is contained in:
@@ -0,0 +1,67 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/google/gar/proto"
|
||||
"github.com/spf13/cobra"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
)
|
||||
|
||||
var (
|
||||
inspectSessionID string
|
||||
inspectServerAddr string
|
||||
)
|
||||
|
||||
var inspectCmd = &cobra.Command{
|
||||
Use: "inspect",
|
||||
Short: "Inspect a session",
|
||||
Long: `Inspect a session to view its current state, step count, and other details.`,
|
||||
RunE: runInspect,
|
||||
}
|
||||
|
||||
func init() {
|
||||
inspectCmd.Flags().StringVar(&inspectSessionID, "session-id", "", "Session ID (required)")
|
||||
inspectCmd.Flags().StringVar(&inspectServerAddr, "server", "", "gRPC controller server address (e.g., localhost:8494)")
|
||||
inspectCmd.MarkFlagRequired("session-id")
|
||||
inspectCmd.MarkFlagRequired("server")
|
||||
}
|
||||
|
||||
func runInspect(cmd *cobra.Command, args []string) error {
|
||||
fmt.Printf("Inspecting session: %s\n", inspectSessionID)
|
||||
|
||||
// Connect to gRPC server
|
||||
conn, err := grpc.NewClient(inspectServerAddr, grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect to server: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
client := proto.NewGARServiceClient(conn)
|
||||
|
||||
// Get session details
|
||||
resp, err := client.GetSession(context.Background(), &proto.GetSessionRequest{
|
||||
SessionId: inspectSessionID,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("error getting session: %w", err)
|
||||
}
|
||||
|
||||
session := resp.Session
|
||||
|
||||
// Print session details
|
||||
fmt.Println("\nSession Details:")
|
||||
fmt.Printf(" ID: %s\n", session.SessionId)
|
||||
fmt.Printf(" State: %s\n", session.State)
|
||||
fmt.Printf(" Current Step: %d\n", session.CurrentStep)
|
||||
fmt.Printf(" Created At: %s\n", time.UnixMilli(session.CreatedAt).Format(time.RFC3339))
|
||||
fmt.Printf(" Updated At: %s\n", time.UnixMilli(session.UpdatedAt).Format(time.RFC3339))
|
||||
fmt.Printf(" Message Count: %d\n", session.MessageCount)
|
||||
fmt.Printf(" Checkpoints: %d\n", session.CheckpointCount)
|
||||
fmt.Printf(" Active Agents: %v\n", session.ActiveAgents)
|
||||
|
||||
return nil
|
||||
}
|
||||
+2
-344
@@ -4,355 +4,13 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
|
||||
"github.com/google/gar/internal/config"
|
||||
"github.com/google/gar/internal/controller"
|
||||
"github.com/google/gar/internal/eventlog"
|
||||
"github.com/google/gar/internal/server"
|
||||
"github.com/google/gar/proto"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
var (
|
||||
command string
|
||||
sessionID string
|
||||
input string
|
||||
checkpointID string
|
||||
agentID string
|
||||
agentName string
|
||||
agentDescription string
|
||||
agentAddr string
|
||||
configFile string // YAML config file path for serve command
|
||||
serverAddr string // gRPC controller server address
|
||||
)
|
||||
|
||||
func main() {
|
||||
// Define subcommands
|
||||
triggerCmd := flag.NewFlagSet("trigger", flag.ExitOnError)
|
||||
inspectCmd := flag.NewFlagSet("inspect", flag.ExitOnError)
|
||||
registerCmd := flag.NewFlagSet("register", flag.ExitOnError)
|
||||
serveCmd := flag.NewFlagSet("serve", flag.ExitOnError)
|
||||
|
||||
// Trigger command flags
|
||||
triggerCmd.StringVar(&sessionID, "session-id", "", "Session ID (optional, generates UUID if not provided, or resumes existing)")
|
||||
triggerCmd.StringVar(&input, "input", "", "Input message to send")
|
||||
triggerCmd.StringVar(&checkpointID, "checkpoint", "", "Resume from specific checkpoint UUID (empty for latest)")
|
||||
triggerCmd.StringVar(&serverAddr, "server", "", "gRPC controller server address (e.g., localhost:8494)")
|
||||
|
||||
// Inspect command flags
|
||||
inspectCmd.StringVar(&sessionID, "session-id", "", "Session ID (required)")
|
||||
inspectCmd.StringVar(&serverAddr, "server", "", "gRPC controller server address (e.g., localhost:8494)")
|
||||
|
||||
// Register command flags
|
||||
registerCmd.StringVar(&agentID, "agent-id", "", "Agent ID (required)")
|
||||
registerCmd.StringVar(&agentName, "name", "", "Agent name")
|
||||
registerCmd.StringVar(&agentDescription, "description", "", "Agent description")
|
||||
registerCmd.StringVar(&agentAddr, "agent-addr", "", "Agent address (e.g., localhost:50051)")
|
||||
registerCmd.StringVar(&serverAddr, "server", "", "gRPC controller server address (e.g., localhost:8494)")
|
||||
|
||||
// Serve command flags
|
||||
serveCmd.StringVar(&configFile, "config", "gar.yaml", "Path to YAML configuration file")
|
||||
|
||||
if len(os.Args) < 2 {
|
||||
printUsage()
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
command = os.Args[1]
|
||||
|
||||
switch command {
|
||||
case "trigger":
|
||||
triggerCmd.Parse(os.Args[2:])
|
||||
runTrigger()
|
||||
|
||||
case "inspect":
|
||||
inspectCmd.Parse(os.Args[2:])
|
||||
if sessionID == "" {
|
||||
fmt.Println("Error: --session-id is required")
|
||||
inspectCmd.Usage()
|
||||
os.Exit(1)
|
||||
}
|
||||
runInspect()
|
||||
|
||||
case "register":
|
||||
registerCmd.Parse(os.Args[2:])
|
||||
if agentID == "" || agentAddr == "" {
|
||||
fmt.Println("Error: --agent-id and --agent-addr are required")
|
||||
registerCmd.Usage()
|
||||
os.Exit(1)
|
||||
}
|
||||
runRegister()
|
||||
|
||||
case "serve":
|
||||
serveCmd.Parse(os.Args[2:])
|
||||
runServe()
|
||||
|
||||
default:
|
||||
printUsage()
|
||||
if err := Execute(); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "Error: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
// connectToServer creates a gRPC client connection to the controller server.
|
||||
func connectToServer() (proto.GARServiceClient, *grpc.ClientConn, error) {
|
||||
conn, err := grpc.NewClient(serverAddr, grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to connect to server: %w", err)
|
||||
}
|
||||
client := proto.NewGARServiceClient(conn)
|
||||
return client, conn, nil
|
||||
}
|
||||
|
||||
func printUsage() {
|
||||
fmt.Println("Usage: gar <command> [options]")
|
||||
fmt.Println()
|
||||
fmt.Println("Commands:")
|
||||
fmt.Println(" trigger Trigger a new session or resume an existing one")
|
||||
fmt.Println(" inspect Inspect a session")
|
||||
fmt.Println(" register Register a remote agent")
|
||||
fmt.Println(" serve Run controller as a gRPC server")
|
||||
fmt.Println()
|
||||
fmt.Println("Use 'gar <command> -h' for more information about a command.")
|
||||
}
|
||||
|
||||
func runTrigger() {
|
||||
// Require server address
|
||||
if serverAddr == "" {
|
||||
fmt.Println("Error: --server flag is required")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Generate UUID if no session ID provided
|
||||
if sessionID == "" {
|
||||
sessionID = uuid.New().String()
|
||||
fmt.Printf("Generated session ID: %s\n", sessionID)
|
||||
}
|
||||
|
||||
fmt.Printf("Triggering session: %s\n", sessionID)
|
||||
|
||||
// Create input content
|
||||
var inputs []*proto.Content
|
||||
if input != "" {
|
||||
inputs = []*proto.Content{
|
||||
{
|
||||
Role: "user",
|
||||
Type: "text",
|
||||
Mimetype: "text/plain",
|
||||
Data: input,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Setup signal handling for graceful shutdown
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, os.Interrupt, syscall.SIGTERM)
|
||||
|
||||
go func() {
|
||||
<-sigChan
|
||||
fmt.Println("\nReceived interrupt, shutting down...")
|
||||
cancel()
|
||||
}()
|
||||
|
||||
// Connect to gRPC server
|
||||
client, conn, err := connectToServer()
|
||||
if err != nil {
|
||||
fmt.Printf("Error connecting to server: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
stream, err := client.TriggerSession(ctx, &proto.TriggerSessionRequest{
|
||||
SessionId: sessionID,
|
||||
Inputs: inputs,
|
||||
CheckpointId: checkpointID,
|
||||
})
|
||||
if err != nil {
|
||||
fmt.Printf("Error triggering session: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Receive and print all responses
|
||||
for {
|
||||
resp, err := stream.Recv()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
fmt.Printf("Error receiving response: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
if resp.Output != nil {
|
||||
fmt.Printf("[%s] %s\n", resp.State, resp.Output.Data)
|
||||
}
|
||||
}
|
||||
|
||||
fmt.Println("Session completed successfully")
|
||||
}
|
||||
|
||||
func runInspect() {
|
||||
// Require server address
|
||||
if serverAddr == "" {
|
||||
fmt.Println("Error: --server flag is required")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("Inspecting session: %s\n", sessionID)
|
||||
|
||||
// Connect to gRPC server
|
||||
client, conn, err := connectToServer()
|
||||
if err != nil {
|
||||
fmt.Printf("Error connecting to server: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Get session details
|
||||
resp, err := client.GetSession(context.Background(), &proto.GetSessionRequest{
|
||||
SessionId: sessionID,
|
||||
})
|
||||
if err != nil {
|
||||
fmt.Printf("Error getting session: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
session := resp.Session
|
||||
|
||||
// Print session details
|
||||
fmt.Println("\nSession Details:")
|
||||
fmt.Printf(" ID: %s\n", session.SessionId)
|
||||
fmt.Printf(" State: %s\n", session.State)
|
||||
fmt.Printf(" Current Step: %d\n", session.CurrentStep)
|
||||
fmt.Printf(" Created At: %s\n", time.UnixMilli(session.CreatedAt).Format(time.RFC3339))
|
||||
fmt.Printf(" Updated At: %s\n", time.UnixMilli(session.UpdatedAt).Format(time.RFC3339))
|
||||
fmt.Printf(" Message Count: %d\n", session.MessageCount)
|
||||
fmt.Printf(" Checkpoints: %d\n", session.CheckpointCount)
|
||||
fmt.Printf(" Active Agents: %v\n", session.ActiveAgents)
|
||||
}
|
||||
|
||||
func runRegister() {
|
||||
// Require server address
|
||||
if serverAddr == "" {
|
||||
fmt.Println("Error: --server flag is required")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("Registering agent: %s at %s\n", agentID, agentAddr)
|
||||
if agentName != "" {
|
||||
fmt.Printf(" Name: %s\n", agentName)
|
||||
}
|
||||
if agentDescription != "" {
|
||||
fmt.Printf(" Description: %s\n", agentDescription)
|
||||
}
|
||||
|
||||
// Connect to gRPC server
|
||||
client, conn, err := connectToServer()
|
||||
if err != nil {
|
||||
fmt.Printf("Error connecting to server: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Register remote agent
|
||||
_, err = client.RegisterAgent(context.Background(), &proto.RegisterAgentRequest{
|
||||
AgentId: agentID,
|
||||
AgentType: "remote",
|
||||
Name: agentName,
|
||||
Description: agentDescription,
|
||||
Address: agentAddr,
|
||||
})
|
||||
if err != nil {
|
||||
fmt.Printf("Error registering agent: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Println("Agent registered successfully")
|
||||
}
|
||||
|
||||
func runServe() {
|
||||
// Load configuration from YAML file
|
||||
cfg, err := config.LoadFromFile(configFile)
|
||||
if err != nil {
|
||||
fmt.Printf("Error loading config file '%s': %v\n", configFile, err)
|
||||
fmt.Println("\nTip: Create a config file with 'gar serve --help' to see an example")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Validate configuration
|
||||
if err := cfg.Validate(); err != nil {
|
||||
fmt.Printf("Invalid configuration: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Printf("Starting GAR server at %s...\n", cfg.Server.Address)
|
||||
fmt.Printf("Event log directory: %s\n", cfg.EventLog.Dir)
|
||||
fmt.Printf("Max steps: %d\n", cfg.Controller.MaxSteps)
|
||||
fmt.Printf("Health check interval: %s\n", cfg.Controller.HealthCheckInterval)
|
||||
|
||||
// Create controller with config
|
||||
c, err := newControllerFromConfig(cfg)
|
||||
if err != nil {
|
||||
fmt.Printf("Error creating controller: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer c.Close()
|
||||
|
||||
// Create server
|
||||
srv := server.New(c)
|
||||
|
||||
// Setup signal handling for graceful shutdown
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, os.Interrupt, syscall.SIGTERM)
|
||||
|
||||
go func() {
|
||||
<-sigChan
|
||||
fmt.Println("\nReceived interrupt, shutting down...")
|
||||
os.Exit(0)
|
||||
}()
|
||||
|
||||
// Start serving
|
||||
if err := srv.Serve(cfg.Server.Address); err != nil {
|
||||
fmt.Printf("Error serving: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func newControllerFromConfig(cfg *config.Config) (*controller.Controller, error) {
|
||||
// Create event log factory
|
||||
eventLogFactory := func(sessionID string) (eventlog.EventLog, error) {
|
||||
return eventlog.NewFileEventLog(eventlog.FileConfig{
|
||||
SessionID: sessionID,
|
||||
Dir: cfg.EventLog.Dir,
|
||||
})
|
||||
}
|
||||
|
||||
// Build controller config
|
||||
// The controller will create a default Gemini planner if PlanFunc is nil
|
||||
// Gemini config can be customized via environment variables (GEMINI_API_KEY, GAR_GEMINI_MODEL)
|
||||
controllerConfig := controller.Config{
|
||||
EventLogFactory: eventLogFactory,
|
||||
MaxSteps: cfg.Controller.MaxSteps,
|
||||
HealthCheckInterval: cfg.Controller.HealthCheckInterval,
|
||||
}
|
||||
|
||||
// Create controller
|
||||
c, err := controller.New(controllerConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/google/gar/proto"
|
||||
"github.com/spf13/cobra"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
)
|
||||
|
||||
var (
|
||||
registerAgentID string
|
||||
registerAgentName string
|
||||
registerAgentDesc string
|
||||
registerAgentAddr string
|
||||
registerServerAddr string
|
||||
)
|
||||
|
||||
var registerCmd = &cobra.Command{
|
||||
Use: "register",
|
||||
Short: "Register a remote agent",
|
||||
Long: `Register a remote agent with the controller so it can be used in sessions.`,
|
||||
RunE: runRegister,
|
||||
}
|
||||
|
||||
func init() {
|
||||
registerCmd.Flags().StringVar(®isterAgentID, "agent-id", "", "Agent ID (required)")
|
||||
registerCmd.Flags().StringVar(®isterAgentName, "name", "", "Agent name")
|
||||
registerCmd.Flags().StringVar(®isterAgentDesc, "description", "", "Agent description")
|
||||
registerCmd.Flags().StringVar(®isterAgentAddr, "agent-addr", "", "Agent address (e.g., localhost:50051)")
|
||||
registerCmd.Flags().StringVar(®isterServerAddr, "server", "", "gRPC controller server address (e.g., localhost:8494)")
|
||||
registerCmd.MarkFlagRequired("agent-id")
|
||||
registerCmd.MarkFlagRequired("agent-addr")
|
||||
registerCmd.MarkFlagRequired("server")
|
||||
}
|
||||
|
||||
func runRegister(cmd *cobra.Command, args []string) error {
|
||||
fmt.Printf("Registering agent: %s at %s\n", registerAgentID, registerAgentAddr)
|
||||
if registerAgentName != "" {
|
||||
fmt.Printf(" Name: %s\n", registerAgentName)
|
||||
}
|
||||
if registerAgentDesc != "" {
|
||||
fmt.Printf(" Description: %s\n", registerAgentDesc)
|
||||
}
|
||||
|
||||
// Connect to gRPC server
|
||||
conn, err := grpc.NewClient(registerServerAddr, grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect to server: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
client := proto.NewGARServiceClient(conn)
|
||||
|
||||
// Register remote agent
|
||||
_, err = client.RegisterAgent(context.Background(), &proto.RegisterAgentRequest{
|
||||
AgentId: registerAgentID,
|
||||
AgentType: "remote",
|
||||
Name: registerAgentName,
|
||||
Description: registerAgentDesc,
|
||||
Address: registerAgentAddr,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("error registering agent: %w", err)
|
||||
}
|
||||
|
||||
fmt.Println("Agent registered successfully")
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
Use: "gar",
|
||||
Short: "GAR - Google Agent Runtime CLI",
|
||||
Long: `Gar is a CLI tool for managing agent orchestrator sessions.
|
||||
It provides commands to trigger sessions, resume from checkpoints,
|
||||
inspect session state, register agents, and run the controller server.`,
|
||||
}
|
||||
|
||||
// Execute runs the root command
|
||||
func Execute() error {
|
||||
return rootCmd.Execute()
|
||||
}
|
||||
|
||||
func init() {
|
||||
// Add subcommands
|
||||
rootCmd.AddCommand(triggerCmd)
|
||||
rootCmd.AddCommand(inspectCmd)
|
||||
rootCmd.AddCommand(registerCmd)
|
||||
rootCmd.AddCommand(serveCmd)
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
|
||||
"github.com/google/gar/internal/config"
|
||||
"github.com/google/gar/internal/controller"
|
||||
"github.com/google/gar/internal/eventlog"
|
||||
"github.com/google/gar/internal/server"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var (
|
||||
serveConfigFile string
|
||||
)
|
||||
|
||||
var serveCmd = &cobra.Command{
|
||||
Use: "serve",
|
||||
Short: "Run controller as a gRPC server",
|
||||
Long: `Run the GAR controller as a gRPC server.
|
||||
Loads configuration from a YAML file (default: gar.yaml).`,
|
||||
RunE: runServe,
|
||||
}
|
||||
|
||||
func init() {
|
||||
serveCmd.Flags().StringVar(&serveConfigFile, "config", "gar.yaml", "Path to YAML configuration file")
|
||||
}
|
||||
|
||||
func runServe(cmd *cobra.Command, args []string) error {
|
||||
// Load configuration from YAML file
|
||||
cfg, err := config.LoadFromFile(serveConfigFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error loading config file '%s': %w\nTip: Create a config file with 'gar serve --help' to see an example", serveConfigFile, err)
|
||||
}
|
||||
|
||||
// Validate configuration
|
||||
if err := cfg.Validate(); err != nil {
|
||||
return fmt.Errorf("invalid configuration: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("Starting GAR server at %s...\n", cfg.Server.Address)
|
||||
fmt.Printf("Event log directory: %s\n", cfg.EventLog.Dir)
|
||||
fmt.Printf("Max steps: %d\n", cfg.Controller.MaxSteps)
|
||||
fmt.Printf("Health check interval: %s\n", cfg.Controller.HealthCheckInterval)
|
||||
|
||||
// Create controller with config
|
||||
c, err := newControllerFromConfig(cfg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error creating controller: %w", err)
|
||||
}
|
||||
defer c.Close()
|
||||
|
||||
// Create server
|
||||
srv := server.New(c)
|
||||
|
||||
// Setup signal handling for graceful shutdown
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, os.Interrupt, syscall.SIGTERM)
|
||||
|
||||
go func() {
|
||||
<-sigChan
|
||||
fmt.Println("\nReceived interrupt, shutting down...")
|
||||
os.Exit(0)
|
||||
}()
|
||||
|
||||
// Start serving
|
||||
if err := srv.Serve(cfg.Server.Address); err != nil {
|
||||
return fmt.Errorf("error serving: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func newControllerFromConfig(cfg *config.Config) (*controller.Controller, error) {
|
||||
// Create event log factory
|
||||
eventLogFactory := func(sessionID string) (eventlog.EventLog, error) {
|
||||
return eventlog.NewFileEventLog(eventlog.FileConfig{
|
||||
SessionID: sessionID,
|
||||
Dir: cfg.EventLog.Dir,
|
||||
})
|
||||
}
|
||||
|
||||
// Build controller config
|
||||
// The controller will create a default Gemini planner if PlanFunc is nil
|
||||
// Gemini config can be customized via environment variables (GEMINI_API_KEY, GAR_GEMINI_MODEL)
|
||||
controllerConfig := controller.Config{
|
||||
EventLogFactory: eventLogFactory,
|
||||
MaxSteps: cfg.Controller.MaxSteps,
|
||||
HealthCheckInterval: cfg.Controller.HealthCheckInterval,
|
||||
}
|
||||
|
||||
// Create controller
|
||||
c, err := controller.New(controllerConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
|
||||
"github.com/google/gar/proto"
|
||||
"github.com/google/uuid"
|
||||
"github.com/spf13/cobra"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
)
|
||||
|
||||
var (
|
||||
triggerSessionID string
|
||||
triggerInput string
|
||||
triggerCheckpoint string
|
||||
triggerServerAddr string
|
||||
)
|
||||
|
||||
var triggerCmd = &cobra.Command{
|
||||
Use: "trigger",
|
||||
Short: "Trigger a new session or resume an existing one",
|
||||
Long: `Trigger a new agentic session or resume an existing one.
|
||||
If no session ID is provided, a new UUID will be generated.
|
||||
Use --checkpoint to resume from a specific checkpoint.`,
|
||||
RunE: runTrigger,
|
||||
}
|
||||
|
||||
func init() {
|
||||
triggerCmd.Flags().StringVar(&triggerSessionID, "session-id", "", "Session ID (optional, generates UUID if not provided)")
|
||||
triggerCmd.Flags().StringVar(&triggerInput, "input", "", "Input message to send")
|
||||
triggerCmd.Flags().StringVar(&triggerCheckpoint, "checkpoint", "", "Resume from specific checkpoint UUID (empty for latest)")
|
||||
triggerCmd.Flags().StringVar(&triggerServerAddr, "server", "", "gRPC controller server address (e.g., localhost:8494)")
|
||||
triggerCmd.MarkFlagRequired("server")
|
||||
}
|
||||
|
||||
func runTrigger(cmd *cobra.Command, args []string) error {
|
||||
// Generate UUID if no session ID provided
|
||||
if triggerSessionID == "" {
|
||||
triggerSessionID = uuid.New().String()
|
||||
fmt.Printf("Generated session ID: %s\n", triggerSessionID)
|
||||
}
|
||||
|
||||
fmt.Printf("Triggering session: %s\n", triggerSessionID)
|
||||
|
||||
// Create input content
|
||||
var inputs []*proto.Content
|
||||
if triggerInput != "" {
|
||||
inputs = []*proto.Content{
|
||||
{
|
||||
Role: "user",
|
||||
Type: "text",
|
||||
Mimetype: "text/plain",
|
||||
Data: triggerInput,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Setup signal handling for graceful shutdown
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, os.Interrupt, syscall.SIGTERM)
|
||||
|
||||
go func() {
|
||||
<-sigChan
|
||||
fmt.Println("\nReceived interrupt, shutting down...")
|
||||
cancel()
|
||||
}()
|
||||
|
||||
// Connect to gRPC server
|
||||
conn, err := grpc.NewClient(triggerServerAddr, grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect to server: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
client := proto.NewGARServiceClient(conn)
|
||||
|
||||
stream, err := client.TriggerSession(ctx, &proto.TriggerSessionRequest{
|
||||
SessionId: triggerSessionID,
|
||||
Inputs: inputs,
|
||||
CheckpointId: triggerCheckpoint,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("error triggering session: %w", err)
|
||||
}
|
||||
|
||||
// Receive and print all responses
|
||||
for {
|
||||
resp, err := stream.Recv()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("error receiving response: %w", err)
|
||||
}
|
||||
|
||||
if resp.Output != nil {
|
||||
fmt.Printf("[%s] %s\n", resp.State, resp.Output.Data)
|
||||
}
|
||||
}
|
||||
|
||||
fmt.Println("Session completed successfully")
|
||||
return nil
|
||||
}
|
||||
@@ -24,6 +24,9 @@ require (
|
||||
github.com/google/s2a-go v0.1.7 // indirect
|
||||
github.com/googleapis/enterprise-certificate-proxy v0.3.2 // indirect
|
||||
github.com/googleapis/gax-go/v2 v2.12.5 // indirect
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/spf13/cobra v1.10.2 // indirect
|
||||
github.com/spf13/pflag v1.0.10 // indirect
|
||||
go.opencensus.io v0.24.0 // indirect
|
||||
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
||||
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.51.0 // indirect
|
||||
|
||||
@@ -15,6 +15,7 @@ github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03
|
||||
github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU=
|
||||
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
|
||||
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
|
||||
@@ -63,8 +64,16 @@ github.com/googleapis/enterprise-certificate-proxy v0.3.2 h1:Vie5ybvEvT75RniqhfF
|
||||
github.com/googleapis/enterprise-certificate-proxy v0.3.2/go.mod h1:VLSiSSBs/ksPL8kq3OBOQ6WRI2QnaFynd1DCjZ62+V0=
|
||||
github.com/googleapis/gax-go/v2 v2.12.5 h1:8gw9KZK8TiVKB6q3zHY3SBzLnrGp6HQjyfYBYGmXdxA=
|
||||
github.com/googleapis/gax-go/v2 v2.12.5/go.mod h1:BUDKcWo+RaKq5SC9vVYL0wLADa3VcfswbOMMRmB9H3E=
|
||||
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
|
||||
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
|
||||
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
|
||||
github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU=
|
||||
github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4=
|
||||
github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk=
|
||||
github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
|
||||
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
|
||||
@@ -89,6 +98,7 @@ go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6
|
||||
go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA=
|
||||
go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE=
|
||||
go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs=
|
||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
|
||||
|
||||
Reference in New Issue
Block a user