mirror of
https://github.com/google/ax.git
synced 2026-10-02 03:14:37 +08:00
Move Harness implementations to their own packages (#243)
* Move Harness implementations to their own packages Fixes #232. * Shorten the constructor names
This commit is contained in:
@@ -23,6 +23,8 @@ import (
|
||||
"github.com/google/ax/internal/controller"
|
||||
"github.com/google/ax/internal/controller/eventlog"
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/internal/harness/antigravity"
|
||||
"github.com/google/ax/internal/harness/substrate"
|
||||
)
|
||||
|
||||
const antigravityHarnessID = "antigravity"
|
||||
@@ -63,9 +65,9 @@ func NewControllerFromConfig(ctx context.Context, cfg *Config) (*controller.Cont
|
||||
if address == "" {
|
||||
address = "127.0.0.1:50053"
|
||||
}
|
||||
antigravityHarness = harness.NewAntigravityHarness(address)
|
||||
antigravityHarness = antigravity.New(address)
|
||||
} else {
|
||||
antigravityHarness, err = harness.NewSubstrateHarness(antigravityHarnessID, "", "", "", 80)
|
||||
antigravityHarness, err = substrate.New(antigravityHarnessID, "", "", "", 80)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("antigravity harness: %w", err)
|
||||
}
|
||||
|
||||
@@ -37,7 +37,7 @@ import (
|
||||
"github.com/google/ax/internal/controller"
|
||||
"github.com/google/ax/internal/controller/eventlog"
|
||||
"github.com/google/ax/internal/controller/eventlog/eventlogtest"
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/internal/harness/antigravity"
|
||||
"github.com/google/ax/proto"
|
||||
)
|
||||
|
||||
@@ -74,8 +74,8 @@ func main() {
|
||||
}
|
||||
conn.Close()
|
||||
fmt.Printf("Connected to Antigravity gRPC harness server at %s\n", address)
|
||||
harness := harness.NewAntigravityHarness(address)
|
||||
reg.RegisterHarness("antigravity", harness)
|
||||
h := antigravity.New(address)
|
||||
reg.RegisterHarness("antigravity", h)
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
"os"
|
||||
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/internal/harness/substrate"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
@@ -110,7 +111,7 @@ func (c SubstrateHarnessConfig) NewHarness(endpoint string) (harness.Harness, er
|
||||
|
||||
// newSubstrateHarness brings up a harness that is deployed as a substrate actor.
|
||||
func newSubstrateHarness(harnessID, endpoint, namespace, template string, port int) (harness.Harness, error) {
|
||||
sh, err := harness.NewSubstrateHarness(harnessID, endpoint, namespace, template, port)
|
||||
sh, err := substrate.New(harnessID, endpoint, namespace, template, port)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package antigravity
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -24,13 +24,14 @@ import (
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/proto"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// Compile-time interface assertions.
|
||||
var _ Harness = (*AntigravityHarness)(nil)
|
||||
var _ Execution = (*antigravityExecution)(nil)
|
||||
var _ harness.Harness = (*AntigravityHarness)(nil)
|
||||
var _ harness.Execution = (*antigravityExecution)(nil)
|
||||
|
||||
// AntigravityHarness implements the Harness interface by connecting to the
|
||||
// Antigravity Python agent server over gRPC.
|
||||
@@ -38,9 +39,9 @@ type AntigravityHarness struct {
|
||||
address string
|
||||
}
|
||||
|
||||
// NewAntigravityHarness creates a new AntigravityHarness with a configurable address.
|
||||
// New creates a new AntigravityHarness with a configurable address.
|
||||
// Address defaults to "127.0.0.1:50053" (gRPC TCP connection).
|
||||
func NewAntigravityHarness(address string) *AntigravityHarness {
|
||||
func New(address string) *AntigravityHarness {
|
||||
if address == "" {
|
||||
address = "127.0.0.1:50053"
|
||||
}
|
||||
@@ -50,7 +51,7 @@ func NewAntigravityHarness(address string) *AntigravityHarness {
|
||||
}
|
||||
|
||||
// Start implements Harness.Start.
|
||||
func (h *AntigravityHarness) Start(ctx context.Context, conversationID string, harnessConfig []byte) (Execution, error) {
|
||||
func (h *AntigravityHarness) Start(ctx context.Context, conversationID string, harnessConfig []byte) (harness.Execution, error) {
|
||||
return &antigravityExecution{
|
||||
harness: h,
|
||||
conversationID: conversationID,
|
||||
@@ -88,7 +89,7 @@ func (e *antigravityExecution) Queue(ctx context.Context, msg ...*proto.Message)
|
||||
}
|
||||
|
||||
// Run executes the turn over gRPC bidirectional streaming and forwards events to the handler.
|
||||
func (e *antigravityExecution) Run(ctx context.Context, handler Handler) error {
|
||||
func (e *antigravityExecution) Run(ctx context.Context, handler harness.Handler) error {
|
||||
ctx, span := otel.Tracer("antigravity-harness").Start(ctx, "Run")
|
||||
defer span.End()
|
||||
|
||||
@@ -143,7 +144,7 @@ func (e *antigravityExecution) Run(ctx context.Context, handler Handler) error {
|
||||
}
|
||||
|
||||
// 5. Stream responses and trigger callbacks
|
||||
return drainStream(ctx, stream, e.id, handler)
|
||||
return harness.DrainStream(ctx, stream, e.id, handler)
|
||||
}
|
||||
|
||||
// Close implements Execution.Close.
|
||||
+14
-13
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package antigravity
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
@@ -20,16 +20,17 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/ax/internal/harness/harnesstest"
|
||||
"github.com/google/ax/proto"
|
||||
)
|
||||
|
||||
var antigravityHarnessConfig = []byte(`{"system_instructions":"be terse"}`)
|
||||
|
||||
func TestAntigravityHarness_Run_Success(t *testing.T) {
|
||||
srv := &mockHarnessServer{
|
||||
outputs: []*proto.Message{thoughtText("Analyzing"), assistantText("Hello world")},
|
||||
srv := &harnesstest.MockHarnessServer{
|
||||
Outputs: []*proto.Message{harnesstest.ThoughtText("Analyzing"), harnesstest.AssistantText("Hello world")},
|
||||
}
|
||||
harnessClient := NewAntigravityHarness(startHarnessServer(t, srv))
|
||||
harnessClient := New(harnesstest.StartHarnessServer(t, srv))
|
||||
|
||||
exec, err := harnessClient.Start(context.Background(), "conv-test", antigravityHarnessConfig)
|
||||
if err != nil {
|
||||
@@ -37,19 +38,19 @@ func TestAntigravityHarness_Run_Success(t *testing.T) {
|
||||
}
|
||||
defer exec.Close(context.Background())
|
||||
|
||||
if err := exec.Queue(context.Background(), userText("Hi")); err != nil {
|
||||
if err := exec.Queue(context.Background(), harnesstest.UserText("Hi")); err != nil {
|
||||
t.Fatalf("failed to queue message: %v", err)
|
||||
}
|
||||
|
||||
handler := &mockHandler{}
|
||||
handler := &harnesstest.MockHandler{}
|
||||
if err := exec.Run(context.Background(), handler); err != nil {
|
||||
t.Fatalf("Run failed: %v", err)
|
||||
}
|
||||
|
||||
if !handler.isDone() {
|
||||
if !handler.IsDone() {
|
||||
t.Error("expected OnComplete to be called")
|
||||
}
|
||||
msgs := handler.collected()
|
||||
msgs := handler.Collected()
|
||||
if len(msgs) != 2 {
|
||||
t.Fatalf("expected 2 messages, got %d", len(msgs))
|
||||
}
|
||||
@@ -60,7 +61,7 @@ func TestAntigravityHarness_Run_Success(t *testing.T) {
|
||||
t.Errorf("expected 'Hello world', got %q", got)
|
||||
}
|
||||
// The harness propagated the conversation id and config to the server.
|
||||
convID, _, harnessConfig, _ := srv.received()
|
||||
convID, _, harnessConfig, _ := srv.Received()
|
||||
if convID != "conv-test" {
|
||||
t.Errorf("server got convID=%q, want conv-test", convID)
|
||||
}
|
||||
@@ -70,17 +71,17 @@ func TestAntigravityHarness_Run_Success(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAntigravityHarness_Run_ErrorFrame(t *testing.T) {
|
||||
srv := &mockHarnessServer{failConnect: true, errMessage: "internal mock server crash"}
|
||||
harnessClient := NewAntigravityHarness(startHarnessServer(t, srv))
|
||||
srv := &harnesstest.MockHarnessServer{FailConnect: true, ErrMessage: "internal mock server crash"}
|
||||
harnessClient := New(harnesstest.StartHarnessServer(t, srv))
|
||||
|
||||
exec, _ := harnessClient.Start(context.Background(), "conv-test", antigravityHarnessConfig)
|
||||
defer exec.Close(context.Background())
|
||||
|
||||
if err := exec.Queue(context.Background(), userText("Hi")); err != nil {
|
||||
if err := exec.Queue(context.Background(), harnesstest.UserText("Hi")); err != nil {
|
||||
t.Fatalf("failed to queue message: %v", err)
|
||||
}
|
||||
|
||||
err := exec.Run(context.Background(), &mockHandler{})
|
||||
err := exec.Run(context.Background(), &harnesstest.MockHandler{})
|
||||
if err == nil {
|
||||
t.Fatal("expected error from Run(), got nil")
|
||||
}
|
||||
+11
-10
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package antigravityinteractions
|
||||
|
||||
// AntigravityInteractionsHarness drives an Antigravity agent through the Vertex
|
||||
// GenAI Interactions API over HTTPS + Server-Sent Events, using the steps-based
|
||||
@@ -58,6 +58,7 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/proto"
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/oauth2"
|
||||
@@ -68,8 +69,8 @@ import (
|
||||
const cloudPlatformScope = "https://www.googleapis.com/auth/cloud-platform"
|
||||
|
||||
// Compile-time interface assertions.
|
||||
var _ Harness = (*AntigravityInteractionsHarness)(nil)
|
||||
var _ Execution = (*antigravityInteractionsExecution)(nil)
|
||||
var _ harness.Harness = (*AntigravityInteractionsHarness)(nil)
|
||||
var _ harness.Execution = (*antigravityInteractionsExecution)(nil)
|
||||
|
||||
const interactionsAPIVersion = "v1beta1"
|
||||
|
||||
@@ -87,7 +88,7 @@ const (
|
||||
const defaultLocation = "global"
|
||||
|
||||
// AntigravityInteractionsConfig configures an AntigravityInteractionsHarness.
|
||||
// Use NewAntigravityInteractionsHarness, which fills sensible defaults.
|
||||
// Use New, which fills sensible defaults.
|
||||
//
|
||||
// Cloud project and location come from the standard GOOGLE_CLOUD_PROJECT and
|
||||
// GOOGLE_CLOUD_LOCATION environment variables.
|
||||
@@ -123,7 +124,7 @@ type AntigravityInteractionsConfig struct {
|
||||
|
||||
// StateDir is the directory where each conversation's resume cursor is
|
||||
// persisted, so a conversation can resume after a restart. It is required:
|
||||
// NewAntigravityInteractionsHarness returns an error if it is empty.
|
||||
// New returns an error if it is empty.
|
||||
//
|
||||
// Correctness relies on a single writer per conversation: writes are
|
||||
// last-write-wins with no compare-and-swap. This is an expectation the caller
|
||||
@@ -172,11 +173,11 @@ type AntigravityInteractionsHarness struct {
|
||||
tsErr error
|
||||
}
|
||||
|
||||
// NewAntigravityInteractionsHarness creates a harness from the given config,
|
||||
// New creates a harness from the given config,
|
||||
// filling in defaults for unset fields. It returns an error if cfg.StateDir is
|
||||
// empty or the cursor store cannot be created: resume-cursor persistence is
|
||||
// required, so a usable state directory must be provided.
|
||||
func NewAntigravityInteractionsHarness(cfg AntigravityInteractionsConfig) (*AntigravityInteractionsHarness, error) {
|
||||
func New(cfg AntigravityInteractionsConfig) (*AntigravityInteractionsHarness, error) {
|
||||
cfg.withDefaults()
|
||||
hc := cfg.HTTPClient
|
||||
if hc == nil {
|
||||
@@ -195,7 +196,7 @@ func NewAntigravityInteractionsHarness(cfg AntigravityInteractionsConfig) (*Anti
|
||||
// Start implements Harness.Start. It loads any previously persisted resume
|
||||
// cursor for conversationID so the returned Execution resumes the existing
|
||||
// interaction chain instead of starting a new one.
|
||||
func (h *AntigravityInteractionsHarness) Start(ctx context.Context, conversationID string, harnessConfig []byte) (Execution, error) {
|
||||
func (h *AntigravityInteractionsHarness) Start(ctx context.Context, conversationID string, harnessConfig []byte) (harness.Execution, error) {
|
||||
e := &antigravityInteractionsExecution{
|
||||
harness: h,
|
||||
conversationID: conversationID,
|
||||
@@ -301,7 +302,7 @@ func (e *antigravityInteractionsExecution) setPrevID(ctx context.Context, id str
|
||||
// The agent's text output is forwarded via handler.OnMessage; handler.OnComplete
|
||||
// is called once when the conversation finishes. At each interaction gap, any
|
||||
// human input queued via Queue (steering) is folded into the next turn.
|
||||
func (e *antigravityInteractionsExecution) Run(ctx context.Context, handler Handler) error {
|
||||
func (e *antigravityInteractionsExecution) Run(ctx context.Context, handler harness.Handler) error {
|
||||
e.mu.Lock()
|
||||
if e.closed {
|
||||
e.mu.Unlock()
|
||||
@@ -383,7 +384,7 @@ func (e *antigravityInteractionsExecution) Run(ctx context.Context, handler Hand
|
||||
}
|
||||
|
||||
// emitText forwards non-empty model text to the handler as a Message.
|
||||
func emitText(ctx context.Context, handler Handler, execID, text string) error {
|
||||
func emitText(ctx context.Context, handler harness.Handler, execID, text string) error {
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return nil
|
||||
}
|
||||
+8
-7
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package antigravityinteractions
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -24,6 +24,7 @@ import (
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/google/ax/internal/harness/harnesstest"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
@@ -89,14 +90,14 @@ func (f *fakeInteractions) recorded() []interactionRequest {
|
||||
func newTestHarness(t *testing.T, fake *fakeInteractions, stateDir string) *AntigravityInteractionsHarness {
|
||||
t.Helper()
|
||||
t.Setenv(envCloudProject, "test-project")
|
||||
h, err := NewAntigravityInteractionsHarness(AntigravityInteractionsConfig{
|
||||
h, err := New(AntigravityInteractionsConfig{
|
||||
Agent: "test-agent",
|
||||
StateDir: stateDir,
|
||||
HTTPClient: &http.Client{Transport: fake},
|
||||
TokenSource: oauth2.StaticTokenSource(&oauth2.Token{AccessToken: "fake-token"}),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewAntigravityInteractionsHarness: %v", err)
|
||||
t.Fatalf("New: %v", err)
|
||||
}
|
||||
return h
|
||||
}
|
||||
@@ -110,10 +111,10 @@ func runOneTurn(t *testing.T, h *AntigravityInteractionsHarness, conversationID,
|
||||
if err != nil {
|
||||
t.Fatalf("Start(%q): %v", conversationID, err)
|
||||
}
|
||||
if err := exec.Queue(ctx, userText(prompt)); err != nil {
|
||||
if err := exec.Queue(ctx, harnesstest.UserText(prompt)); err != nil {
|
||||
t.Fatalf("Queue: %v", err)
|
||||
}
|
||||
if err := exec.Run(ctx, &mockHandler{}); err != nil {
|
||||
if err := exec.Run(ctx, &harnesstest.MockHandler{}); err != nil {
|
||||
t.Fatalf("Run: %v", err)
|
||||
}
|
||||
if err := exec.Close(ctx); err != nil {
|
||||
@@ -155,13 +156,13 @@ func TestResumeAcrossRestart(t *testing.T) {
|
||||
// StateDir: resume-cursor persistence is required.
|
||||
func TestNewRequiresStateDir(t *testing.T) {
|
||||
t.Setenv(envCloudProject, "test-project")
|
||||
_, err := NewAntigravityInteractionsHarness(AntigravityInteractionsConfig{
|
||||
_, err := New(AntigravityInteractionsConfig{
|
||||
Agent: "test-agent",
|
||||
StateDir: "", // missing
|
||||
TokenSource: oauth2.StaticTokenSource(&oauth2.Token{AccessToken: "fake-token"}),
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("NewAntigravityInteractionsHarness with empty StateDir: got nil error, want error")
|
||||
t.Fatal("New with empty StateDir: got nil error, want error")
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package antigravityinteractions
|
||||
|
||||
import (
|
||||
"context"
|
||||
+1
-1
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package antigravityinteractions
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package harnesstest
|
||||
|
||||
// Shared in-process mocks for the harness tests: a mock Substrate Control server
|
||||
// (the substrate control plane), a mock HarnessService server (the harness
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/agent-substrate/substrate/pkg/proto/ateapipb"
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/proto"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
@@ -37,7 +38,7 @@ import (
|
||||
// actor lifecycle calls SubstrateHarness makes and lets tests steer the
|
||||
// CreateActor/ResumeActor responses. Only the three RPCs SubstrateHarness uses
|
||||
// are implemented; the rest come from the embedded Unimplemented server.
|
||||
type mockControlServer struct {
|
||||
type MockControlServer struct {
|
||||
ateapipb.UnimplementedControlServer
|
||||
|
||||
mu sync.Mutex
|
||||
@@ -45,44 +46,44 @@ type mockControlServer struct {
|
||||
resumeCalls []string
|
||||
suspendCalls []string
|
||||
|
||||
createErr error // returned from CreateActor when non-nil
|
||||
resumeIP string // AteomPodIp returned from ResumeActor
|
||||
resumeNilActor bool // when true, ResumeActor returns a nil Actor
|
||||
CreateErr error // returned from CreateActor when non-nil
|
||||
ResumeIP string // AteomPodIp returned from ResumeActor
|
||||
ResumeNilActor bool // when true, ResumeActor returns a nil Actor
|
||||
}
|
||||
|
||||
func (f *mockControlServer) CreateAtespace(_ context.Context, req *ateapipb.CreateAtespaceRequest) (*ateapipb.CreateAtespaceResponse, error) {
|
||||
func (f *MockControlServer) CreateAtespace(_ context.Context, req *ateapipb.CreateAtespaceRequest) (*ateapipb.CreateAtespaceResponse, error) {
|
||||
return &ateapipb.CreateAtespaceResponse{Atespace: &ateapipb.Atespace{Name: req.GetName()}}, nil
|
||||
}
|
||||
|
||||
func (f *mockControlServer) CreateActor(_ context.Context, req *ateapipb.CreateActorRequest) (*ateapipb.CreateActorResponse, error) {
|
||||
func (f *MockControlServer) CreateActor(_ context.Context, req *ateapipb.CreateActorRequest) (*ateapipb.CreateActorResponse, error) {
|
||||
f.mu.Lock()
|
||||
f.createCalls = append(f.createCalls, req.GetActorRef().GetName())
|
||||
f.mu.Unlock()
|
||||
if f.createErr != nil {
|
||||
return nil, f.createErr
|
||||
if f.CreateErr != nil {
|
||||
return nil, f.CreateErr
|
||||
}
|
||||
return &ateapipb.CreateActorResponse{Actor: &ateapipb.Actor{ActorId: req.GetActorRef().GetName()}}, nil
|
||||
}
|
||||
|
||||
func (f *mockControlServer) ResumeActor(_ context.Context, req *ateapipb.ResumeActorRequest) (*ateapipb.ResumeActorResponse, error) {
|
||||
func (f *MockControlServer) ResumeActor(_ context.Context, req *ateapipb.ResumeActorRequest) (*ateapipb.ResumeActorResponse, error) {
|
||||
f.mu.Lock()
|
||||
f.resumeCalls = append(f.resumeCalls, req.GetActorRef().GetName())
|
||||
f.mu.Unlock()
|
||||
if f.resumeNilActor {
|
||||
if f.ResumeNilActor {
|
||||
return &ateapipb.ResumeActorResponse{}, nil
|
||||
}
|
||||
return &ateapipb.ResumeActorResponse{Actor: &ateapipb.Actor{ActorId: req.GetActorRef().GetName(), AteomPodIp: f.resumeIP}}, nil
|
||||
return &ateapipb.ResumeActorResponse{Actor: &ateapipb.Actor{ActorId: req.GetActorRef().GetName(), AteomPodIp: f.ResumeIP}}, nil
|
||||
}
|
||||
|
||||
func (f *mockControlServer) SuspendActor(_ context.Context, req *ateapipb.SuspendActorRequest) (*ateapipb.SuspendActorResponse, error) {
|
||||
func (f *MockControlServer) SuspendActor(_ context.Context, req *ateapipb.SuspendActorRequest) (*ateapipb.SuspendActorResponse, error) {
|
||||
f.mu.Lock()
|
||||
f.suspendCalls = append(f.suspendCalls, req.GetActorRef().GetName())
|
||||
f.mu.Unlock()
|
||||
return &ateapipb.SuspendActorResponse{}, nil
|
||||
}
|
||||
|
||||
// calls returns copies of the recorded call lists.
|
||||
func (f *mockControlServer) calls() (create, resume, suspend []string) {
|
||||
// Calls returns copies of the recorded call lists.
|
||||
func (f *MockControlServer) Calls() (create, resume, suspend []string) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return append([]string(nil), f.createCalls...),
|
||||
@@ -94,20 +95,20 @@ func (f *mockControlServer) calls() (create, resume, suspend []string) {
|
||||
// the harness running inside an actor (substrate) or a local subprocess
|
||||
// (antigravity). It records the start frame and emits its configured outputs
|
||||
// followed by a terminal HarnessEnd.
|
||||
type mockHarnessServer struct {
|
||||
type MockHarnessServer struct {
|
||||
proto.UnimplementedHarnessServiceServer
|
||||
|
||||
// outputs are the messages emitted (in a single Outputs frame) before the
|
||||
// Outputs are the messages emitted (in a single Outputs frame) before the
|
||||
// terminal HarnessEnd. When nil, each input is echoed as "ack: <input>".
|
||||
outputs []*proto.Message
|
||||
// failConnect makes Connect return an RPC error before any frame.
|
||||
failConnect bool
|
||||
// failFrame makes Connect terminate the turn with HarnessEnd{STATE_FAILED}.
|
||||
failFrame bool
|
||||
// errCode is the error code used by failFrame.
|
||||
errCode int32
|
||||
// errMessage is the error text used by failConnect/failFrame.
|
||||
errMessage string
|
||||
Outputs []*proto.Message
|
||||
// FailConnect makes Connect return an RPC error before any frame.
|
||||
FailConnect bool
|
||||
// FailFrame makes Connect terminate the turn with HarnessEnd{STATE_FAILED}.
|
||||
FailFrame bool
|
||||
// ErrCode is the error code used by FailFrame.
|
||||
ErrCode int32
|
||||
// ErrMessage is the error text used by FailConnect/FailFrame.
|
||||
ErrMessage string
|
||||
|
||||
mu sync.Mutex
|
||||
gotConvID string
|
||||
@@ -116,9 +117,9 @@ type mockHarnessServer struct {
|
||||
gotInputs []string
|
||||
}
|
||||
|
||||
func (s *mockHarnessServer) Connect(stream proto.HarnessService_ConnectServer) error {
|
||||
if s.failConnect {
|
||||
return status.Error(codes.Internal, s.errMessage)
|
||||
func (s *MockHarnessServer) Connect(stream proto.HarnessService_ConnectServer) error {
|
||||
if s.FailConnect {
|
||||
return status.Error(codes.Internal, s.ErrMessage)
|
||||
}
|
||||
|
||||
req, err := stream.Recv()
|
||||
@@ -140,25 +141,25 @@ func (s *mockHarnessServer) Connect(stream proto.HarnessService_ConnectServer) e
|
||||
s.mu.Unlock()
|
||||
|
||||
convID := req.GetConversationId()
|
||||
if s.failFrame {
|
||||
if s.FailFrame {
|
||||
return stream.Send(&proto.HarnessResponse{
|
||||
ConversationId: convID,
|
||||
Type: &proto.HarnessResponse_End{
|
||||
End: &proto.HarnessEnd{
|
||||
State: proto.State_STATE_FAILED,
|
||||
Error: &proto.Error{
|
||||
Code: s.errCode,
|
||||
Description: s.errMessage,
|
||||
Code: s.ErrCode,
|
||||
Description: s.ErrMessage,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
msgs := s.outputs
|
||||
msgs := s.Outputs
|
||||
if msgs == nil {
|
||||
for _, in := range inputs {
|
||||
msgs = append(msgs, assistantText("ack: "+in))
|
||||
msgs = append(msgs, AssistantText("ack: "+in))
|
||||
}
|
||||
}
|
||||
if len(msgs) > 0 {
|
||||
@@ -177,49 +178,51 @@ func (s *mockHarnessServer) Connect(stream proto.HarnessService_ConnectServer) e
|
||||
})
|
||||
}
|
||||
|
||||
// received returns a copy of the start frame the server received.
|
||||
func (s *mockHarnessServer) received() (convID, harnessID string, harnessConfig []byte, inputs []string) {
|
||||
// Received returns a copy of the start frame the server received.
|
||||
func (s *MockHarnessServer) Received() (convID, harnessID string, harnessConfig []byte, inputs []string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.gotConvID, s.gotHarnessID, append([]byte(nil), s.gotHarnessConfig...), append([]string(nil), s.gotInputs...)
|
||||
}
|
||||
|
||||
// mockHandler records the messages and completion streamed during a turn.
|
||||
type mockHandler struct {
|
||||
type MockHandler struct {
|
||||
mu sync.Mutex
|
||||
messages []*proto.Message
|
||||
complete bool
|
||||
}
|
||||
|
||||
func (h *mockHandler) OnMessage(_ context.Context, _ string, msg *proto.Message) error {
|
||||
var _ harness.Handler = (*MockHandler)(nil)
|
||||
|
||||
func (h *MockHandler) OnMessage(_ context.Context, _ string, msg *proto.Message) error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
h.messages = append(h.messages, msg)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *mockHandler) OnComplete(_ context.Context, _ string) error {
|
||||
func (h *MockHandler) OnComplete(_ context.Context, _ string) error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
h.complete = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *mockHandler) isDone() bool {
|
||||
func (h *MockHandler) IsDone() bool {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
return h.complete
|
||||
}
|
||||
|
||||
// collected returns a copy of the messages received via OnMessage.
|
||||
func (h *mockHandler) collected() []*proto.Message {
|
||||
// Collected returns a copy of the messages received via OnMessage.
|
||||
func (h *MockHandler) Collected() []*proto.Message {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
return append([]*proto.Message(nil), h.messages...)
|
||||
}
|
||||
|
||||
// texts returns the text content of each received message, in order.
|
||||
func (h *mockHandler) texts() []string {
|
||||
// Texts returns the text content of each received message, in order.
|
||||
func (h *MockHandler) Texts() []string {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
var out []string
|
||||
@@ -229,21 +232,21 @@ func (h *mockHandler) texts() []string {
|
||||
return out
|
||||
}
|
||||
|
||||
func assistantText(text string) *proto.Message {
|
||||
func AssistantText(text string) *proto.Message {
|
||||
return &proto.Message{
|
||||
Role: "assistant",
|
||||
Content: &proto.Content{Type: &proto.Content_Text{Text: &proto.TextContent{Text: text}}},
|
||||
}
|
||||
}
|
||||
|
||||
func userText(text string) *proto.Message {
|
||||
func UserText(text string) *proto.Message {
|
||||
return &proto.Message{
|
||||
Role: "user",
|
||||
Content: &proto.Content{Type: &proto.Content_Text{Text: &proto.TextContent{Text: text}}},
|
||||
}
|
||||
}
|
||||
|
||||
func thoughtText(summary string) *proto.Message {
|
||||
func ThoughtText(summary string) *proto.Message {
|
||||
return &proto.Message{
|
||||
Role: "model",
|
||||
Content: &proto.Content{
|
||||
@@ -258,9 +261,9 @@ func thoughtText(summary string) *proto.Message {
|
||||
}
|
||||
}
|
||||
|
||||
// startHarnessServer starts a HarnessService + health server (status SERVING)
|
||||
// StartHarnessServer starts a HarnessService + health server (status SERVING)
|
||||
// on a random local port and returns its address.
|
||||
func startHarnessServer(t *testing.T, srv *mockHarnessServer) string {
|
||||
func StartHarnessServer(t *testing.T, srv *MockHarnessServer) string {
|
||||
t.Helper()
|
||||
lis, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
@@ -276,8 +279,8 @@ func startHarnessServer(t *testing.T, srv *mockHarnessServer) string {
|
||||
return lis.Addr().String()
|
||||
}
|
||||
|
||||
// startControlServer starts a mock Substrate Control server on a random local port.
|
||||
func startControlServer(t *testing.T, srv *mockControlServer) string {
|
||||
// StartControlServer starts a mock Substrate Control server on a random local port.
|
||||
func StartControlServer(t *testing.T, srv *MockControlServer) string {
|
||||
t.Helper()
|
||||
lis, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
@@ -22,9 +22,9 @@ import (
|
||||
"github.com/google/ax/proto"
|
||||
)
|
||||
|
||||
// drainStream reads from the harness gRPC stream until io.EOF, dispatching messages
|
||||
// DrainStream reads from the harness gRPC stream until io.EOF, dispatching messages
|
||||
// to the handler, and returns the final execution status.
|
||||
func drainStream(ctx context.Context, stream proto.HarnessService_ConnectClient, execID string, handler Handler) error {
|
||||
func DrainStream(ctx context.Context, stream proto.HarnessService_ConnectClient, execID string, handler Handler) error {
|
||||
var endState proto.State
|
||||
var endErr error
|
||||
hasEnd := false
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package substrate
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -32,11 +32,16 @@ import (
|
||||
"google.golang.org/grpc/health/grpc_health_v1"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/google/ax/internal/harness"
|
||||
"github.com/google/ax/internal/k8s/ate"
|
||||
"github.com/google/ax/proto"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// Compile-time interface assertions.
|
||||
var _ harness.Harness = (*SubstrateHarness)(nil)
|
||||
var _ harness.Execution = (*substrateExecution)(nil)
|
||||
|
||||
// healthCheckTimeout defines the maximum time Start waits for a freshly
|
||||
// created/resumed actor's harness to become reachable and ready.
|
||||
const healthCheckTimeout = 60 * time.Second
|
||||
@@ -49,8 +54,8 @@ type SubstrateHarness struct {
|
||||
dialOpts []grpc.DialOption
|
||||
}
|
||||
|
||||
// NewSubstrateHarness creates a new SubstrateHarness.
|
||||
func NewSubstrateHarness(harnessID string, endpoint string, namespace string, template string, port int, opts ...grpc.DialOption) (*SubstrateHarness, error) {
|
||||
// New creates a new SubstrateHarness.
|
||||
func New(harnessID string, endpoint string, namespace string, template string, port int, opts ...grpc.DialOption) (*SubstrateHarness, error) {
|
||||
if port == 0 {
|
||||
port = 50053 // Default HarnessService port
|
||||
}
|
||||
@@ -78,7 +83,7 @@ func NewSubstrateHarness(harnessID string, endpoint string, namespace string, te
|
||||
}
|
||||
|
||||
// Start implements Harness interface. It creates/resumes the target actor.
|
||||
func (h *SubstrateHarness) Start(ctx context.Context, conversationID string, harnessConfig []byte) (Execution, error) {
|
||||
func (h *SubstrateHarness) Start(ctx context.Context, conversationID string, harnessConfig []byte) (harness.Execution, error) {
|
||||
if conversationID == "" {
|
||||
return nil, errors.New("SubstrateHarness needs valid conversationID")
|
||||
}
|
||||
@@ -185,7 +190,7 @@ func (e *substrateExecution) Queue(ctx context.Context, msg ...*proto.Message) e
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *substrateExecution) Run(ctx context.Context, handler Handler) error {
|
||||
func (e *substrateExecution) Run(ctx context.Context, handler harness.Handler) error {
|
||||
ctx, span := otel.Tracer("substrate-harness").Start(ctx, "Run")
|
||||
defer span.End()
|
||||
|
||||
@@ -220,7 +225,7 @@ func (e *substrateExecution) Run(ctx context.Context, handler Handler) error {
|
||||
}
|
||||
|
||||
// Drain HarnessResponse frames until the terminal HarnessEnd.
|
||||
return drainStream(ctx, stream, e.execID, handler)
|
||||
return harness.DrainStream(ctx, stream, e.execID, handler)
|
||||
}
|
||||
|
||||
func (e *substrateExecution) Close(ctx context.Context) error {
|
||||
@@ -12,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package harness
|
||||
package substrate
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
@@ -24,6 +24,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/ax/internal/harness/harnesstest"
|
||||
"github.com/google/ax/internal/k8s/ate"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
@@ -156,25 +157,25 @@ func newTestSubstrateHarness(t *testing.T, ctrlAddr, harnessAddr string) *Substr
|
||||
// break: create/resume idempotency, worker-IP extraction, the health gate, the
|
||||
// Connect streaming protocol, and suspend-on-close.
|
||||
func TestSubstrateHarness_EndToEnd(t *testing.T) {
|
||||
ctrl := &mockControlServer{resumeIP: "127.0.0.1"}
|
||||
srv := &mockHarnessServer{}
|
||||
h := newTestSubstrateHarness(t, startControlServer(t, ctrl), startHarnessServer(t, srv))
|
||||
ctrl := &harnesstest.MockControlServer{ResumeIP: "127.0.0.1"}
|
||||
srv := &harnesstest.MockHarnessServer{}
|
||||
h := newTestSubstrateHarness(t, harnesstest.StartControlServer(t, ctrl), harnesstest.StartHarnessServer(t, srv))
|
||||
|
||||
ctx := context.Background()
|
||||
exec, err := h.Start(ctx, "conv-1", substrateHarnessConfig)
|
||||
if err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
if err := exec.Queue(ctx, userText("hi")); err != nil {
|
||||
if err := exec.Queue(ctx, harnesstest.UserText("hi")); err != nil {
|
||||
t.Fatalf("Queue: %v", err)
|
||||
}
|
||||
handler := &mockHandler{}
|
||||
handler := &harnesstest.MockHandler{}
|
||||
if err := exec.Run(ctx, handler); err != nil {
|
||||
t.Fatalf("Run: %v", err)
|
||||
}
|
||||
|
||||
// The harness server received the start frame with the right identifiers.
|
||||
convID, harnessID, harnessConfig, inputs := srv.received()
|
||||
convID, harnessID, harnessConfig, inputs := srv.Received()
|
||||
if convID != "conv-1" || harnessID != "antigravity" {
|
||||
t.Errorf("server got convID=%q harnessID=%q, want conv-1/antigravity", convID, harnessID)
|
||||
}
|
||||
@@ -186,15 +187,15 @@ func TestSubstrateHarness_EndToEnd(t *testing.T) {
|
||||
}
|
||||
|
||||
// The handler streamed the output and completed.
|
||||
if !handler.isDone() {
|
||||
if !handler.IsDone() {
|
||||
t.Error("handler did not complete")
|
||||
}
|
||||
if got := handler.texts(); !slices.Equal(got, []string{"ack: hi"}) {
|
||||
if got := handler.Texts(); !slices.Equal(got, []string{"ack: hi"}) {
|
||||
t.Errorf("handler messages=%v, want [ack: hi]", got)
|
||||
}
|
||||
|
||||
// CreateActor then ResumeActor ran for the conversation; no suspend yet.
|
||||
create, resume, suspend := ctrl.calls()
|
||||
create, resume, suspend := ctrl.Calls()
|
||||
want := []string{"conv-1"}
|
||||
if !slices.Equal(create, want) || !slices.Equal(resume, want) {
|
||||
t.Errorf("create=%v resume=%v, want %v each", create, resume, want)
|
||||
@@ -207,17 +208,17 @@ func TestSubstrateHarness_EndToEnd(t *testing.T) {
|
||||
if err := exec.Close(ctx); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if _, _, suspend = ctrl.calls(); !slices.Equal(suspend, want) {
|
||||
if _, _, suspend = ctrl.Calls(); !slices.Equal(suspend, want) {
|
||||
t.Errorf("suspend=%v, want %v", suspend, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSubstrateHarness_CreateAlreadyExistsTolerated(t *testing.T) {
|
||||
ctrl := &mockControlServer{
|
||||
resumeIP: "127.0.0.1",
|
||||
createErr: status.Error(codes.AlreadyExists, "exists"),
|
||||
ctrl := &harnesstest.MockControlServer{
|
||||
ResumeIP: "127.0.0.1",
|
||||
CreateErr: status.Error(codes.AlreadyExists, "exists"),
|
||||
}
|
||||
h := newTestSubstrateHarness(t, startControlServer(t, ctrl), startHarnessServer(t, &mockHarnessServer{}))
|
||||
h := newTestSubstrateHarness(t, harnesstest.StartControlServer(t, ctrl), harnesstest.StartHarnessServer(t, &harnesstest.MockHarnessServer{}))
|
||||
|
||||
ctx := context.Background()
|
||||
exec, err := h.Start(ctx, "conv-1", substrateHarnessConfig)
|
||||
@@ -226,24 +227,24 @@ func TestSubstrateHarness_CreateAlreadyExistsTolerated(t *testing.T) {
|
||||
}
|
||||
t.Cleanup(func() { _ = exec.Close(ctx) })
|
||||
|
||||
if err := exec.Queue(ctx, userText("hi")); err != nil {
|
||||
if err := exec.Queue(ctx, harnesstest.UserText("hi")); err != nil {
|
||||
t.Fatalf("Queue: %v", err)
|
||||
}
|
||||
handler := &mockHandler{}
|
||||
handler := &harnesstest.MockHandler{}
|
||||
if err := exec.Run(ctx, handler); err != nil {
|
||||
t.Fatalf("Run: %v", err)
|
||||
}
|
||||
if !handler.isDone() {
|
||||
if !handler.IsDone() {
|
||||
t.Error("handler did not complete")
|
||||
}
|
||||
if _, resume, _ := ctrl.calls(); !slices.Equal(resume, []string{"conv-1"}) {
|
||||
if _, resume, _ := ctrl.Calls(); !slices.Equal(resume, []string{"conv-1"}) {
|
||||
t.Errorf("resume=%v, want [conv-1]", resume)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSubstrateHarness_ResumeNoWorkerIP(t *testing.T) {
|
||||
ctrl := &mockControlServer{resumeIP: ""} // empty AteomPodIp
|
||||
h := newTestSubstrateHarness(t, startControlServer(t, ctrl), startHarnessServer(t, &mockHarnessServer{}))
|
||||
ctrl := &harnesstest.MockControlServer{ResumeIP: ""} // empty AteomPodIp
|
||||
h := newTestSubstrateHarness(t, harnesstest.StartControlServer(t, ctrl), harnesstest.StartHarnessServer(t, &harnesstest.MockHarnessServer{}))
|
||||
|
||||
_, err := h.Start(context.Background(), "conv-1", substrateHarnessConfig)
|
||||
if err == nil {
|
||||
@@ -255,8 +256,8 @@ func TestSubstrateHarness_ResumeNoWorkerIP(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSubstrateHarness_ResumeNilActor(t *testing.T) {
|
||||
ctrl := &mockControlServer{resumeNilActor: true}
|
||||
h := newTestSubstrateHarness(t, startControlServer(t, ctrl), startHarnessServer(t, &mockHarnessServer{}))
|
||||
ctrl := &harnesstest.MockControlServer{ResumeNilActor: true}
|
||||
h := newTestSubstrateHarness(t, harnesstest.StartControlServer(t, ctrl), harnesstest.StartHarnessServer(t, &harnesstest.MockHarnessServer{}))
|
||||
|
||||
_, err := h.Start(context.Background(), "conv-1", substrateHarnessConfig)
|
||||
if err == nil {
|
||||
@@ -268,9 +269,9 @@ func TestSubstrateHarness_ResumeNilActor(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSubstrateHarness_HarnessFailedFrame(t *testing.T) {
|
||||
ctrl := &mockControlServer{resumeIP: "127.0.0.1"}
|
||||
srv := &mockHarnessServer{failFrame: true, errCode: 13, errMessage: "boom"}
|
||||
h := newTestSubstrateHarness(t, startControlServer(t, ctrl), startHarnessServer(t, srv))
|
||||
ctrl := &harnesstest.MockControlServer{ResumeIP: "127.0.0.1"}
|
||||
srv := &harnesstest.MockHarnessServer{FailFrame: true, ErrCode: 13, ErrMessage: "boom"}
|
||||
h := newTestSubstrateHarness(t, harnesstest.StartControlServer(t, ctrl), harnesstest.StartHarnessServer(t, srv))
|
||||
|
||||
ctx := context.Background()
|
||||
exec, err := h.Start(ctx, "conv-1", substrateHarnessConfig)
|
||||
@@ -278,10 +279,10 @@ func TestSubstrateHarness_HarnessFailedFrame(t *testing.T) {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = exec.Close(ctx) })
|
||||
if err := exec.Queue(ctx, userText("hi")); err != nil {
|
||||
if err := exec.Queue(ctx, harnesstest.UserText("hi")); err != nil {
|
||||
t.Fatalf("Queue: %v", err)
|
||||
}
|
||||
if err := exec.Run(ctx, &mockHandler{}); err == nil {
|
||||
if err := exec.Run(ctx, &harnesstest.MockHandler{}); err == nil {
|
||||
t.Fatal("expected error from failed harness frame, got nil")
|
||||
} else if !strings.Contains(err.Error(), "harness failed") {
|
||||
t.Errorf("error = %v, want it to mention 'harness failed'", err)
|
||||
Reference in New Issue
Block a user