diff --git a/agent/agent.pb.go b/agent/agent.pb.go index 2b9d0984..865439e7 100644 --- a/agent/agent.pb.go +++ b/agent/agent.pb.go @@ -441,24 +441,24 @@ var file_agent_agent_proto_rawDesc = []byte{ 0x18, 0x01, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x0a, 0x72, 0x65, 0x70, 0x6f, 0x72, 0x74, 0x44, 0x61, 0x74, 0x61, 0x22, 0x29, 0x0a, 0x13, 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x12, 0x0a, 0x04, 0x66, 0x69, 0x6c, - 0x65, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x04, 0x66, 0x69, 0x6c, 0x65, 0x32, 0xf5, 0x01, - 0x0a, 0x0c, 0x41, 0x67, 0x65, 0x6e, 0x74, 0x53, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x12, 0x31, + 0x65, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x04, 0x66, 0x69, 0x6c, 0x65, 0x32, 0xf9, 0x01, + 0x0a, 0x0c, 0x41, 0x67, 0x65, 0x6e, 0x74, 0x53, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x12, 0x33, 0x0a, 0x04, 0x41, 0x6c, 0x67, 0x6f, 0x12, 0x12, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x41, 0x6c, 0x67, 0x6f, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x13, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x41, 0x6c, 0x67, 0x6f, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x22, - 0x00, 0x12, 0x31, 0x0a, 0x04, 0x44, 0x61, 0x74, 0x61, 0x12, 0x12, 0x2e, 0x61, 0x67, 0x65, 0x6e, - 0x74, 0x2e, 0x44, 0x61, 0x74, 0x61, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x13, 0x2e, - 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x44, 0x61, 0x74, 0x61, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, - 0x73, 0x65, 0x22, 0x00, 0x12, 0x37, 0x0a, 0x06, 0x52, 0x65, 0x73, 0x75, 0x6c, 0x74, 0x12, 0x14, - 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x52, 0x65, 0x73, 0x75, 0x6c, 0x74, 0x52, 0x65, 0x71, - 0x75, 0x65, 0x73, 0x74, 0x1a, 0x15, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x52, 0x65, 0x73, - 0x75, 0x6c, 0x74, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x22, 0x00, 0x12, 0x46, 0x0a, - 0x0b, 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x12, 0x19, 0x2e, 0x61, - 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, - 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x1a, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, - 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, - 0x6e, 0x73, 0x65, 0x22, 0x00, 0x42, 0x09, 0x5a, 0x07, 0x2e, 0x2f, 0x61, 0x67, 0x65, 0x6e, 0x74, - 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33, + 0x00, 0x28, 0x01, 0x12, 0x33, 0x0a, 0x04, 0x44, 0x61, 0x74, 0x61, 0x12, 0x12, 0x2e, 0x61, 0x67, + 0x65, 0x6e, 0x74, 0x2e, 0x44, 0x61, 0x74, 0x61, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, + 0x13, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x44, 0x61, 0x74, 0x61, 0x52, 0x65, 0x73, 0x70, + 0x6f, 0x6e, 0x73, 0x65, 0x22, 0x00, 0x28, 0x01, 0x12, 0x37, 0x0a, 0x06, 0x52, 0x65, 0x73, 0x75, + 0x6c, 0x74, 0x12, 0x14, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x52, 0x65, 0x73, 0x75, 0x6c, + 0x74, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x15, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, + 0x2e, 0x52, 0x65, 0x73, 0x75, 0x6c, 0x74, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x22, + 0x00, 0x12, 0x46, 0x0a, 0x0b, 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, + 0x12, 0x19, 0x2e, 0x61, 0x67, 0x65, 0x6e, 0x74, 0x2e, 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, + 0x74, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x1a, 0x2e, 0x61, 0x67, + 0x65, 0x6e, 0x74, 0x2e, 0x41, 0x74, 0x74, 0x65, 0x73, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x52, + 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x22, 0x00, 0x42, 0x09, 0x5a, 0x07, 0x2e, 0x2f, 0x61, + 0x67, 0x65, 0x6e, 0x74, 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33, } var ( diff --git a/agent/agent.proto b/agent/agent.proto index 721eb6d3..6876324a 100644 --- a/agent/agent.proto +++ b/agent/agent.proto @@ -8,8 +8,8 @@ package agent; option go_package = "./agent"; service AgentService { - rpc Algo(AlgoRequest) returns (AlgoResponse) {} - rpc Data(DataRequest) returns (DataResponse) {} + rpc Algo(stream AlgoRequest) returns (AlgoResponse) {} + rpc Data(stream DataRequest) returns (DataResponse) {} rpc Result(ResultRequest) returns (ResultResponse) {} rpc Attestation(AttestationRequest) returns (AttestationResponse) {} } diff --git a/agent/agent_grpc.pb.go b/agent/agent_grpc.pb.go index 60e3e7d3..4c7d5489 100644 --- a/agent/agent_grpc.pb.go +++ b/agent/agent_grpc.pb.go @@ -32,8 +32,8 @@ const ( // // For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. type AgentServiceClient interface { - Algo(ctx context.Context, in *AlgoRequest, opts ...grpc.CallOption) (*AlgoResponse, error) - Data(ctx context.Context, in *DataRequest, opts ...grpc.CallOption) (*DataResponse, error) + Algo(ctx context.Context, opts ...grpc.CallOption) (AgentService_AlgoClient, error) + Data(ctx context.Context, opts ...grpc.CallOption) (AgentService_DataClient, error) Result(ctx context.Context, in *ResultRequest, opts ...grpc.CallOption) (*ResultResponse, error) Attestation(ctx context.Context, in *AttestationRequest, opts ...grpc.CallOption) (*AttestationResponse, error) } @@ -46,22 +46,72 @@ func NewAgentServiceClient(cc grpc.ClientConnInterface) AgentServiceClient { return &agentServiceClient{cc} } -func (c *agentServiceClient) Algo(ctx context.Context, in *AlgoRequest, opts ...grpc.CallOption) (*AlgoResponse, error) { - out := new(AlgoResponse) - err := c.cc.Invoke(ctx, AgentService_Algo_FullMethodName, in, out, opts...) +func (c *agentServiceClient) Algo(ctx context.Context, opts ...grpc.CallOption) (AgentService_AlgoClient, error) { + stream, err := c.cc.NewStream(ctx, &AgentService_ServiceDesc.Streams[0], AgentService_Algo_FullMethodName, opts...) if err != nil { return nil, err } - return out, nil + x := &agentServiceAlgoClient{stream} + return x, nil } -func (c *agentServiceClient) Data(ctx context.Context, in *DataRequest, opts ...grpc.CallOption) (*DataResponse, error) { - out := new(DataResponse) - err := c.cc.Invoke(ctx, AgentService_Data_FullMethodName, in, out, opts...) +type AgentService_AlgoClient interface { + Send(*AlgoRequest) error + CloseAndRecv() (*AlgoResponse, error) + grpc.ClientStream +} + +type agentServiceAlgoClient struct { + grpc.ClientStream +} + +func (x *agentServiceAlgoClient) Send(m *AlgoRequest) error { + return x.ClientStream.SendMsg(m) +} + +func (x *agentServiceAlgoClient) CloseAndRecv() (*AlgoResponse, error) { + if err := x.ClientStream.CloseSend(); err != nil { + return nil, err + } + m := new(AlgoResponse) + if err := x.ClientStream.RecvMsg(m); err != nil { + return nil, err + } + return m, nil +} + +func (c *agentServiceClient) Data(ctx context.Context, opts ...grpc.CallOption) (AgentService_DataClient, error) { + stream, err := c.cc.NewStream(ctx, &AgentService_ServiceDesc.Streams[1], AgentService_Data_FullMethodName, opts...) if err != nil { return nil, err } - return out, nil + x := &agentServiceDataClient{stream} + return x, nil +} + +type AgentService_DataClient interface { + Send(*DataRequest) error + CloseAndRecv() (*DataResponse, error) + grpc.ClientStream +} + +type agentServiceDataClient struct { + grpc.ClientStream +} + +func (x *agentServiceDataClient) Send(m *DataRequest) error { + return x.ClientStream.SendMsg(m) +} + +func (x *agentServiceDataClient) CloseAndRecv() (*DataResponse, error) { + if err := x.ClientStream.CloseSend(); err != nil { + return nil, err + } + m := new(DataResponse) + if err := x.ClientStream.RecvMsg(m); err != nil { + return nil, err + } + return m, nil } func (c *agentServiceClient) Result(ctx context.Context, in *ResultRequest, opts ...grpc.CallOption) (*ResultResponse, error) { @@ -86,8 +136,8 @@ func (c *agentServiceClient) Attestation(ctx context.Context, in *AttestationReq // All implementations must embed UnimplementedAgentServiceServer // for forward compatibility type AgentServiceServer interface { - Algo(context.Context, *AlgoRequest) (*AlgoResponse, error) - Data(context.Context, *DataRequest) (*DataResponse, error) + Algo(AgentService_AlgoServer) error + Data(AgentService_DataServer) error Result(context.Context, *ResultRequest) (*ResultResponse, error) Attestation(context.Context, *AttestationRequest) (*AttestationResponse, error) mustEmbedUnimplementedAgentServiceServer() @@ -97,11 +147,11 @@ type AgentServiceServer interface { type UnimplementedAgentServiceServer struct { } -func (UnimplementedAgentServiceServer) Algo(context.Context, *AlgoRequest) (*AlgoResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method Algo not implemented") +func (UnimplementedAgentServiceServer) Algo(AgentService_AlgoServer) error { + return status.Errorf(codes.Unimplemented, "method Algo not implemented") } -func (UnimplementedAgentServiceServer) Data(context.Context, *DataRequest) (*DataResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method Data not implemented") +func (UnimplementedAgentServiceServer) Data(AgentService_DataServer) error { + return status.Errorf(codes.Unimplemented, "method Data not implemented") } func (UnimplementedAgentServiceServer) Result(context.Context, *ResultRequest) (*ResultResponse, error) { return nil, status.Errorf(codes.Unimplemented, "method Result not implemented") @@ -122,40 +172,56 @@ func RegisterAgentServiceServer(s grpc.ServiceRegistrar, srv AgentServiceServer) s.RegisterService(&AgentService_ServiceDesc, srv) } -func _AgentService_Algo_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(AlgoRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AgentServiceServer).Algo(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AgentService_Algo_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AgentServiceServer).Algo(ctx, req.(*AlgoRequest)) - } - return interceptor(ctx, in, info, handler) +func _AgentService_Algo_Handler(srv interface{}, stream grpc.ServerStream) error { + return srv.(AgentServiceServer).Algo(&agentServiceAlgoServer{stream}) } -func _AgentService_Data_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(DataRequest) - if err := dec(in); err != nil { +type AgentService_AlgoServer interface { + SendAndClose(*AlgoResponse) error + Recv() (*AlgoRequest, error) + grpc.ServerStream +} + +type agentServiceAlgoServer struct { + grpc.ServerStream +} + +func (x *agentServiceAlgoServer) SendAndClose(m *AlgoResponse) error { + return x.ServerStream.SendMsg(m) +} + +func (x *agentServiceAlgoServer) Recv() (*AlgoRequest, error) { + m := new(AlgoRequest) + if err := x.ServerStream.RecvMsg(m); err != nil { return nil, err } - if interceptor == nil { - return srv.(AgentServiceServer).Data(ctx, in) + return m, nil +} + +func _AgentService_Data_Handler(srv interface{}, stream grpc.ServerStream) error { + return srv.(AgentServiceServer).Data(&agentServiceDataServer{stream}) +} + +type AgentService_DataServer interface { + SendAndClose(*DataResponse) error + Recv() (*DataRequest, error) + grpc.ServerStream +} + +type agentServiceDataServer struct { + grpc.ServerStream +} + +func (x *agentServiceDataServer) SendAndClose(m *DataResponse) error { + return x.ServerStream.SendMsg(m) +} + +func (x *agentServiceDataServer) Recv() (*DataRequest, error) { + m := new(DataRequest) + if err := x.ServerStream.RecvMsg(m); err != nil { + return nil, err } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AgentService_Data_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AgentServiceServer).Data(ctx, req.(*DataRequest)) - } - return interceptor(ctx, in, info, handler) + return m, nil } func _AgentService_Result_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { @@ -201,14 +267,6 @@ var AgentService_ServiceDesc = grpc.ServiceDesc{ ServiceName: "agent.AgentService", HandlerType: (*AgentServiceServer)(nil), Methods: []grpc.MethodDesc{ - { - MethodName: "Algo", - Handler: _AgentService_Algo_Handler, - }, - { - MethodName: "Data", - Handler: _AgentService_Data_Handler, - }, { MethodName: "Result", Handler: _AgentService_Result_Handler, @@ -218,6 +276,17 @@ var AgentService_ServiceDesc = grpc.ServiceDesc{ Handler: _AgentService_Attestation_Handler, }, }, - Streams: []grpc.StreamDesc{}, + Streams: []grpc.StreamDesc{ + { + StreamName: "Algo", + Handler: _AgentService_Algo_Handler, + ClientStreams: true, + }, + { + StreamName: "Data", + Handler: _AgentService_Data_Handler, + ClientStreams: true, + }, + }, Metadata: "agent/agent.proto", } diff --git a/agent/api/grpc/client.go b/agent/api/grpc/client.go deleted file mode 100644 index b35a404b..00000000 --- a/agent/api/grpc/client.go +++ /dev/null @@ -1,218 +0,0 @@ -// Copyright (c) Ultraviolet -// SPDX-License-Identifier: Apache-2.0 -package grpc - -import ( - "context" - "fmt" - "time" - - "github.com/go-kit/kit/endpoint" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "github.com/ultravioletrs/cocos/agent" - "google.golang.org/grpc" -) - -const svcName = "agent.AgentService" - -type grpcClient struct { - algo endpoint.Endpoint - data endpoint.Endpoint - result endpoint.Endpoint - attestation endpoint.Endpoint - timeout time.Duration -} - -// NewClient returns new gRPC client instance. -func NewClient(conn *grpc.ClientConn, timeout time.Duration) agent.AgentServiceClient { - return &grpcClient{ - algo: kitgrpc.NewClient( - conn, - svcName, - "Algo", - encodeAlgoRequest, - decodeAlgoResponse, - agent.AlgoResponse{}, - ).Endpoint(), - data: kitgrpc.NewClient( - conn, - svcName, - "Data", - encodeDataRequest, - decodeDataResponse, - agent.DataResponse{}, - ).Endpoint(), - result: kitgrpc.NewClient( - conn, - svcName, - "Result", - encodeResultRequest, - decodeResultResponse, - agent.ResultResponse{}, - ).Endpoint(), - attestation: kitgrpc.NewClient( - conn, - svcName, - "Attestation", - encodeAttestationRequest, - decodeAttestationResponse, - agent.AttestationResponse{}, - ).Endpoint(), - timeout: timeout, - } -} - -// encodeAlgoRequest is a transport/grpc.EncodeRequestFunc that -// converts a user-domain algoReq to a gRPC request. -func encodeAlgoRequest(_ context.Context, request interface{}) (interface{}, error) { - req, ok := request.(*algoReq) - if !ok { - return nil, fmt.Errorf("invalid request type: %T", request) - } - - return &agent.AlgoRequest{ - Algorithm: req.Algorithm, - Provider: req.Provider, - Id: req.Id, - }, nil -} - -// decodeAlgoResponse is a transport/grpc.DecodeResponseFunc that -// converts a gRPC AlgoResponse to a user-domain response. -func decodeAlgoResponse(_ context.Context, grpcResponse interface{}) (interface{}, error) { - _, ok := grpcResponse.(*agent.AlgoResponse) - if !ok { - return nil, fmt.Errorf("invalid response type: %T", grpcResponse) - } - - return algoRes{}, nil -} - -// encodeDataRequest is a transport/grpc.EncodeRequestFunc that -// converts a user-domain dataReq to a gRPC request. -func encodeDataRequest(_ context.Context, request interface{}) (interface{}, error) { - req, ok := request.(*dataReq) - if !ok { - return nil, fmt.Errorf("invalid request type: %T", request) - } - - return &agent.DataRequest{ - Dataset: req.Dataset, - Provider: req.Provider, - Id: req.Id, - }, nil -} - -// decodeDataResponse is a transport/grpc.DecodeResponseFunc that -// converts a gRPC DataResponse to a user-domain response. -func decodeDataResponse(_ context.Context, grpcResponse interface{}) (interface{}, error) { - _, ok := grpcResponse.(*agent.DataResponse) - if !ok { - return nil, fmt.Errorf("invalid response type: %T", grpcResponse) - } - - return dataRes{}, nil -} - -// encodeResultRequest is a transport/grpc.EncodeRequestFunc that -// converts a user-domain resultReq to a gRPC request. -func encodeResultRequest(_ context.Context, request interface{}) (interface{}, error) { - req, ok := request.(*resultReq) - if !ok { - return nil, fmt.Errorf("invalid request type: %T", request) - } - - return &agent.ResultRequest{ - Consumer: req.Consumer, - }, nil -} - -// decodeResultResponse is a transport/grpc.DecodeResponseFunc that -// converts a gRPC ResultResponse to a user-domain response. -func decodeResultResponse(_ context.Context, grpcResponse interface{}) (interface{}, error) { - response, ok := grpcResponse.(*agent.ResultResponse) - if !ok { - return nil, fmt.Errorf("invalid response type: %T", grpcResponse) - } - - return resultRes{ - File: response.File, - }, nil -} - -// encodeAttestationRequest is a transport/grpc.EncodeRequestFunc that -// converts a user-domain attestationReq to a gRPC request. -func encodeAttestationRequest(_ context.Context, request interface{}) (interface{}, error) { - req, ok := request.(*attestationReq) - if !ok { - return nil, fmt.Errorf("invalid request type: %T", request) - } - return &agent.AttestationRequest{ReportData: req.ReportData[:]}, nil -} - -// decodeAttestationResponse is a transport/grpc.DecodeResponseFunc that -// converts a gRPC AttestationResponse to a user-domain response. -func decodeAttestationResponse(_ context.Context, grpcResponse interface{}) (interface{}, error) { - response, ok := grpcResponse.(*agent.AttestationResponse) - if !ok { - return nil, fmt.Errorf("invalid response type: %T", grpcResponse) - } - - return attestationRes{ - File: response.File, - }, nil -} - -// Algo implements the Algo method of the agent.AgentServiceClient interface. -func (c grpcClient) Algo(ctx context.Context, request *agent.AlgoRequest, _ ...grpc.CallOption) (*agent.AlgoResponse, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - _, err := c.algo(ctx, &algoReq{Algorithm: request.Algorithm, Provider: request.Provider, Id: request.Id}) - if err != nil { - return nil, err - } - - return &agent.AlgoResponse{}, nil -} - -// Data implements the Data method of the agent.AgentServiceClient interface. -func (c grpcClient) Data(ctx context.Context, request *agent.DataRequest, _ ...grpc.CallOption) (*agent.DataResponse, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - _, err := c.data(ctx, &dataReq{Dataset: request.Dataset, Provider: request.Provider, Id: request.Id}) - if err != nil { - return nil, err - } - - return &agent.DataResponse{}, nil -} - -// Result implements the Result method of the agent.AgentServiceClient interface. -func (c grpcClient) Result(ctx context.Context, request *agent.ResultRequest, _ ...grpc.CallOption) (*agent.ResultResponse, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - res, err := c.result(ctx, &resultReq{Consumer: request.Consumer}) - if err != nil { - return nil, err - } - - resultRes := res.(resultRes) - return &agent.ResultResponse{File: resultRes.File}, nil -} - -// Result implements the Result method of the agent.AgentServiceClient interface. -func (c grpcClient) Attestation(ctx context.Context, request *agent.AttestationRequest, _ ...grpc.CallOption) (*agent.AttestationResponse, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - res, err := c.attestation(ctx, &attestationReq{ReportData: [agent.ReportDataSize]byte(request.ReportData)}) - if err != nil { - return nil, err - } - - attestationRes := res.(attestationRes) - return &agent.AttestationResponse{File: attestationRes.File}, nil -} diff --git a/agent/api/grpc/server.go b/agent/api/grpc/server.go index 38295fde..0bfc7c6a 100644 --- a/agent/api/grpc/server.go +++ b/agent/api/grpc/server.go @@ -5,11 +5,16 @@ package grpc import ( "context" "errors" + "io" "github.com/go-kit/kit/transport/grpc" "github.com/ultravioletrs/cocos/agent" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" ) +var _ agent.AgentServiceServer = (*grpcServer)(nil) + type grpcServer struct { algo grpc.Handler data grpc.Handler @@ -99,22 +104,52 @@ func encodeAttestationResponse(_ context.Context, response interface{}) (interfa }, nil } -func (s *grpcServer) Algo(ctx context.Context, req *agent.AlgoRequest) (*agent.AlgoResponse, error) { - _, res, err := s.algo.ServeGRPC(ctx, req) +// Algo implements agent.AgentServiceServer. +func (s *grpcServer) Algo(stream agent.AgentService_AlgoServer) error { + var algoFile []byte + var provider, id string + for { + algoChunk, err := stream.Recv() + if err == io.EOF { + break + } + if err != nil { + return status.Error(codes.Internal, err.Error()) + } + provider = algoChunk.Provider + id = algoChunk.Id + algoFile = append(algoFile, algoChunk.Algorithm...) + } + _, res, err := s.algo.ServeGRPC(stream.Context(), &agent.AlgoRequest{Algorithm: algoFile, Provider: provider, Id: id}) if err != nil { - return nil, err + return err } ar := res.(*agent.AlgoResponse) - return ar, nil + return stream.SendAndClose(ar) } -func (s *grpcServer) Data(ctx context.Context, req *agent.DataRequest) (*agent.DataResponse, error) { - _, res, err := s.data.ServeGRPC(ctx, req) - if err != nil { - return nil, err +// Data implements agent.AgentServiceServer. +func (s *grpcServer) Data(stream agent.AgentService_DataServer) error { + var dataFile []byte + var provider, id string + for { + dataChunk, err := stream.Recv() + if err == io.EOF { + break + } + if err != nil { + return status.Error(codes.Internal, err.Error()) + } + provider = dataChunk.Provider + id = dataChunk.Id + dataFile = append(dataFile, dataChunk.Dataset...) } - dr := res.(*agent.DataResponse) - return dr, nil + _, res, err := s.data.ServeGRPC(stream.Context(), &agent.DataRequest{Dataset: dataFile, Provider: provider, Id: id}) + if err != nil { + return err + } + ar := res.(*agent.DataResponse) + return stream.SendAndClose(ar) } func (s *grpcServer) Result(ctx context.Context, req *agent.ResultRequest) (*agent.ResultResponse, error) { diff --git a/agent/service.go b/agent/service.go index 3a6c5312..cc5f0673 100644 --- a/agent/service.go +++ b/agent/service.go @@ -106,17 +106,16 @@ func (as *agentService) Algo(ctx context.Context, algorithm Algorithm) error { hash := sha3.Sum256(algorithm.Algorithm) - index := containsID(as.computation.Algorithm, algorithm.ID) - switch index { - case -1: + if as.computation.Algorithm.ID != algorithm.ID { return errUndeclaredAlgorithm - default: - if as.computation.Algorithm.Provider != algorithm.Provider { - return errProviderMissmatch - } - if hash != as.computation.Algorithm.Hash { - return errHashMismatch - } + } + + if as.computation.Algorithm.Provider != algorithm.Provider { + return errProviderMissmatch + } + + if hash != as.computation.Algorithm.Hash { + return errHashMismatch } as.algorithm = algorithm.Algorithm diff --git a/pkg/clients/grpc/agent/agent.go b/pkg/clients/grpc/agent/agent.go index 5c1751b6..8e0ebf5f 100644 --- a/pkg/clients/grpc/agent/agent.go +++ b/pkg/clients/grpc/agent/agent.go @@ -4,7 +4,6 @@ package agent import ( "github.com/ultravioletrs/cocos/agent" - agentapi "github.com/ultravioletrs/cocos/agent/api/grpc" "github.com/ultravioletrs/cocos/pkg/clients/grpc" ) @@ -15,5 +14,5 @@ func NewAgentClient(cfg grpc.Config) (grpc.Client, agent.AgentServiceClient, err return nil, nil, err } - return client, agentapi.NewClient(client.Connection(), cfg.Timeout), nil + return client, agent.NewAgentServiceClient(client.Connection()), nil } diff --git a/pkg/sdk/agent.go b/pkg/sdk/agent.go index 95f83747..244d5d8f 100644 --- a/pkg/sdk/agent.go +++ b/pkg/sdk/agent.go @@ -3,7 +3,9 @@ package sdk import ( + "bytes" "context" + "io" "log/slog" "github.com/ultravioletrs/cocos/agent" @@ -11,7 +13,10 @@ import ( var _ agent.Service = (*agentSDK)(nil) -const size64 = 64 +const ( + size64 = 64 + bufferSize = 1024 * 1024 +) type agentSDK struct { client agent.AgentServiceClient @@ -26,14 +31,30 @@ func NewAgentSDK(log *slog.Logger, agentClient agent.AgentServiceClient) *agentS } func (sdk *agentSDK) Algo(ctx context.Context, algorithm agent.Algorithm) error { - request := &agent.AlgoRequest{ - Algorithm: algorithm.Algorithm, - Provider: algorithm.Provider, - Id: algorithm.ID, + stream, err := sdk.client.Algo(ctx) + if err != nil { + sdk.logger.Error("Failed to call Algo RPC") + return err + } + algoBuffer := bytes.NewBuffer(algorithm.Algorithm) + + buf := make([]byte, bufferSize) + for { + n, err := algoBuffer.Read(buf) + if err == io.EOF { + break + } + if err != nil { + return err + } + + err = stream.Send(&agent.AlgoRequest{Id: algorithm.ID, Provider: algorithm.Provider, Algorithm: buf[:n]}) + if err != nil { + return err + } } - if _, err := sdk.client.Algo(ctx, request); err != nil { - sdk.logger.Error("Failed to call Algo RPC") + if _, err := stream.CloseAndRecv(); err != nil { return err } @@ -41,14 +62,30 @@ func (sdk *agentSDK) Algo(ctx context.Context, algorithm agent.Algorithm) error } func (sdk *agentSDK) Data(ctx context.Context, dataset agent.Dataset) error { - request := &agent.DataRequest{ - Dataset: dataset.Dataset, - Provider: dataset.Provider, - Id: dataset.ID, + stream, err := sdk.client.Data(ctx) + if err != nil { + sdk.logger.Error("Failed to call Algo RPC") + return err + } + dataBuffer := bytes.NewBuffer(dataset.Dataset) + + buf := make([]byte, bufferSize) + for { + n, err := dataBuffer.Read(buf) + if err == io.EOF { + break + } + if err != nil { + return err + } + + err = stream.Send(&agent.DataRequest{Id: dataset.ID, Provider: dataset.Provider, Dataset: buf[:n]}) + if err != nil { + return err + } } - if _, err := sdk.client.Data(ctx, request); err != nil { - sdk.logger.Error("Failed to call Data RPC") + if _, err := stream.CloseAndRecv(); err != nil { return err } diff --git a/test/manual/algo/README.md b/test/manual/algo/README.md index 0820199b..f166d559 100644 --- a/test/manual/algo/README.md +++ b/test/manual/algo/README.md @@ -5,5 +5,5 @@ In this example we'll use [pyinstaller](https://pypi.org/project/pyinstaller/) ```shell pip install -U pyinstaller -pyinstaller lin_reg.py +pyinstaller --onefile lin_reg.py ```