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:
Jaana Dogan
2026-07-06 14:00:57 -07:00
committed by GitHub
parent 72b51a4893
commit e39b1232b2
13 changed files with 149 additions and 133 deletions
+4 -2
View File
@@ -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)
}
+3 -3
View File
@@ -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)
})
}
+2 -1
View File
@@ -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.
@@ -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")
}
@@ -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
}
@@ -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")
}
}
@@ -12,7 +12,7 @@
// See the License for the specific language governing permissions and
// limitations under the License.
package harness
package antigravityinteractions
import (
"context"
@@ -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 {
+2 -2
View File
@@ -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)