mirror of
https://github.com/google/ax.git
synced 2026-10-02 03:14:37 +08:00
Move checkpoint_id into Content
This commit is contained in:
@@ -220,20 +220,32 @@ func (sm *SessionManager) CloseAll() {
|
||||
}
|
||||
}
|
||||
|
||||
// WriteContentIn appends an incoming content message to the session with a new checkpoint.
|
||||
// WriteContentIn appends an incoming content message to the session.
|
||||
// Creates a checkpoint only if checkpoint_id is provided in the content.
|
||||
func (s *Session) WriteContentIn(ctx context.Context, content *proto.Content) (string, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
// Generate a new checkpoint UUID
|
||||
checkpointID := uuid.New().String()
|
||||
// Use checkpoint_id from content if provided
|
||||
checkpointID := content.CheckpointId
|
||||
|
||||
if checkpointID != "" {
|
||||
// TODO(jbd): Optimize the lookup.
|
||||
for _, existingID := range s.CheckpointIDs {
|
||||
if existingID == checkpointID {
|
||||
return "", fmt.Errorf("checkpoint %s already exists", checkpointID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if err := s.eventLog.AppendContent(ctx, eventlog.EventTypeContentIn, checkpointID, content); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
s.MessageHistory = append(s.MessageHistory, content)
|
||||
s.CheckpointIDs = append(s.CheckpointIDs, checkpointID)
|
||||
if checkpointID != "" {
|
||||
s.CheckpointIDs = append(s.CheckpointIDs, checkpointID)
|
||||
}
|
||||
s.UpdatedAt = time.Now()
|
||||
return checkpointID, nil
|
||||
}
|
||||
|
||||
@@ -65,16 +65,9 @@ func (s *Server) TriggerSession(req *proto.TriggerSessionRequest, stream grpc.Se
|
||||
return err
|
||||
}
|
||||
|
||||
// Get the latest checkpoint ID if available
|
||||
latestCheckpointID := ""
|
||||
if len(session.CheckpointIDs) > 0 {
|
||||
latestCheckpointID = session.CheckpointIDs[len(session.CheckpointIDs)-1]
|
||||
}
|
||||
|
||||
// Send final success response
|
||||
return stream.Send(&proto.TriggerSessionResponse{
|
||||
State: session.State,
|
||||
CheckpointId: latestCheckpointID,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
+29
-29
@@ -84,10 +84,11 @@ func (State) EnumDescriptor() ([]byte, []int) {
|
||||
// Content represents a message with role, type, mimetype, and data fields
|
||||
type Content struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Role string `protobuf:"bytes,1,opt,name=role,proto3" json:"role,omitempty"` // The role of the content (e.g., "user", "assistant", "system")
|
||||
Type string `protobuf:"bytes,2,opt,name=type,proto3" json:"type,omitempty"` // The type of content (e.g., "text", "image")
|
||||
Mimetype string `protobuf:"bytes,3,opt,name=mimetype,proto3" json:"mimetype,omitempty"` // MIME type of the data (e.g., "text/plain", "image/png")
|
||||
Data string `protobuf:"bytes,4,opt,name=data,proto3" json:"data,omitempty"` // The actual content data
|
||||
Role string `protobuf:"bytes,1,opt,name=role,proto3" json:"role,omitempty"` // The role of the content (e.g., "user", "assistant", "system")
|
||||
Type string `protobuf:"bytes,2,opt,name=type,proto3" json:"type,omitempty"` // The type of content (e.g., "text", "image")
|
||||
Mimetype string `protobuf:"bytes,3,opt,name=mimetype,proto3" json:"mimetype,omitempty"` // MIME type of the data (e.g., "text/plain", "image/png")
|
||||
Data string `protobuf:"bytes,4,opt,name=data,proto3" json:"data,omitempty"` // The actual content data
|
||||
CheckpointId string `protobuf:"bytes,5,opt,name=checkpoint_id,json=checkpointId,proto3" json:"checkpoint_id,omitempty"` // Optional: Checkpoint ID for this content
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -150,6 +151,13 @@ func (x *Content) GetData() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *Content) GetCheckpointId() string {
|
||||
if x != nil {
|
||||
return x.CheckpointId
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// LifecycleEvent represents events from agent lifecycle
|
||||
type LifecycleEvent struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
@@ -305,8 +313,8 @@ func (x *HealthCheckResponse) GetMessage() string {
|
||||
type TriggerSessionRequest struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` // Unique session identifier
|
||||
Inputs []*Content `protobuf:"bytes,2,rep,name=inputs,proto3" json:"inputs,omitempty"` // Input content to process
|
||||
CheckpointId string `protobuf:"bytes,3,opt,name=checkpoint_id,json=checkpointId,proto3" json:"checkpoint_id,omitempty"` // Optional: Resume from specific checkpoint UUID (empty for latest)
|
||||
CheckpointId string `protobuf:"bytes,2,opt,name=checkpoint_id,json=checkpointId,proto3" json:"checkpoint_id,omitempty"` // Optional: Resume from specific checkpoint UUID (empty for latest)
|
||||
Inputs []*Content `protobuf:"bytes,3,rep,name=inputs,proto3" json:"inputs,omitempty"` // Input content to process
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -348,13 +356,6 @@ func (x *TriggerSessionRequest) GetSessionId() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *TriggerSessionRequest) GetInputs() []*Content {
|
||||
if x != nil {
|
||||
return x.Inputs
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *TriggerSessionRequest) GetCheckpointId() string {
|
||||
if x != nil {
|
||||
return x.CheckpointId
|
||||
@@ -362,12 +363,18 @@ func (x *TriggerSessionRequest) GetCheckpointId() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *TriggerSessionRequest) GetInputs() []*Content {
|
||||
if x != nil {
|
||||
return x.Inputs
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// TriggerSessionResponse contains the result of triggering a session
|
||||
type TriggerSessionResponse struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
State State `protobuf:"varint,1,opt,name=state,proto3,enum=proto.State" json:"state,omitempty"` // Session state
|
||||
Output *Content `protobuf:"bytes,2,opt,name=output,proto3" json:"output,omitempty"`
|
||||
CheckpointId string `protobuf:"bytes,3,opt,name=checkpoint_id,json=checkpointId,proto3" json:"checkpoint_id,omitempty"` // Checkpoint UUID for this response
|
||||
Output *Content `protobuf:"bytes,2,opt,name=output,proto3" json:"output,omitempty"` // Output content
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -416,13 +423,6 @@ func (x *TriggerSessionResponse) GetOutput() *Content {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *TriggerSessionResponse) GetCheckpointId() string {
|
||||
if x != nil {
|
||||
return x.CheckpointId
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// GetSessionRequest for retrieving session details
|
||||
type GetSessionRequest struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
@@ -806,12 +806,13 @@ var File_proto_gar_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_proto_gar_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x0fproto/gar.proto\x12\x05proto\x1a\x1fgoogle/protobuf/timestamp.proto\"a\n" +
|
||||
"\x0fproto/gar.proto\x12\x05proto\x1a\x1fgoogle/protobuf/timestamp.proto\"\x86\x01\n" +
|
||||
"\aContent\x12\x12\n" +
|
||||
"\x04role\x18\x01 \x01(\tR\x04role\x12\x12\n" +
|
||||
"\x04type\x18\x02 \x01(\tR\x04type\x12\x1a\n" +
|
||||
"\bmimetype\x18\x03 \x01(\tR\bmimetype\x12\x12\n" +
|
||||
"\x04data\x18\x04 \x01(\tR\x04data\"\xe7\x01\n" +
|
||||
"\x04data\x18\x04 \x01(\tR\x04data\x12#\n" +
|
||||
"\rcheckpoint_id\x18\x05 \x01(\tR\fcheckpointId\"\xe7\x01\n" +
|
||||
"\x0eLifecycleEvent\x12\x1d\n" +
|
||||
"\n" +
|
||||
"event_type\x18\x01 \x01(\tR\teventType\x128\n" +
|
||||
@@ -826,13 +827,12 @@ const file_proto_gar_proto_rawDesc = "" +
|
||||
"\amessage\x18\x02 \x01(\tR\amessage\"\x83\x01\n" +
|
||||
"\x15TriggerSessionRequest\x12\x1d\n" +
|
||||
"\n" +
|
||||
"session_id\x18\x01 \x01(\tR\tsessionId\x12&\n" +
|
||||
"\x06inputs\x18\x02 \x03(\v2\x0e.proto.ContentR\x06inputs\x12#\n" +
|
||||
"\rcheckpoint_id\x18\x03 \x01(\tR\fcheckpointId\"\x89\x01\n" +
|
||||
"session_id\x18\x01 \x01(\tR\tsessionId\x12#\n" +
|
||||
"\rcheckpoint_id\x18\x02 \x01(\tR\fcheckpointId\x12&\n" +
|
||||
"\x06inputs\x18\x03 \x03(\v2\x0e.proto.ContentR\x06inputs\"d\n" +
|
||||
"\x16TriggerSessionResponse\x12\"\n" +
|
||||
"\x05state\x18\x01 \x01(\x0e2\f.proto.StateR\x05state\x12&\n" +
|
||||
"\x06output\x18\x02 \x01(\v2\x0e.proto.ContentR\x06output\x12#\n" +
|
||||
"\rcheckpoint_id\x18\x03 \x01(\tR\fcheckpointId\"2\n" +
|
||||
"\x06output\x18\x02 \x01(\v2\x0e.proto.ContentR\x06output\"2\n" +
|
||||
"\x11GetSessionRequest\x12\x1d\n" +
|
||||
"\n" +
|
||||
"session_id\x18\x01 \x01(\tR\tsessionId\"\xbf\x02\n" +
|
||||
|
||||
+4
-4
@@ -12,6 +12,7 @@ message Content {
|
||||
string type = 2; // The type of content (e.g., "text", "image")
|
||||
string mimetype = 3; // MIME type of the data (e.g., "text/plain", "image/png")
|
||||
string data = 4; // The actual content data
|
||||
string checkpoint_id = 5; // Optional: Checkpoint ID for this content
|
||||
|
||||
// TODO: Replace Content with Interactions Content.
|
||||
}
|
||||
@@ -57,15 +58,14 @@ enum State {
|
||||
// TriggerSessionRequest for triggering a new session
|
||||
message TriggerSessionRequest {
|
||||
string session_id = 1; // Unique session identifier
|
||||
repeated Content inputs = 2; // Input content to process
|
||||
string checkpoint_id = 3; // Optional: Resume from specific checkpoint UUID (empty for latest)
|
||||
string checkpoint_id = 2; // Optional: Resume from specific checkpoint UUID (empty for latest)
|
||||
repeated Content inputs = 3; // Input content to process
|
||||
}
|
||||
|
||||
// TriggerSessionResponse contains the result of triggering a session
|
||||
message TriggerSessionResponse {
|
||||
State state = 1; // Session state
|
||||
Content output = 2;
|
||||
string checkpoint_id = 3; // Checkpoint UUID for this response
|
||||
Content output = 2; // Output content
|
||||
}
|
||||
|
||||
// GetSessionRequest for retrieving session details
|
||||
|
||||
Reference in New Issue
Block a user