diff --git a/.github/workflows/cicd.yaml b/.github/workflows/cicd.yaml index 7130a7b..09318c9 100644 --- a/.github/workflows/cicd.yaml +++ b/.github/workflows/cicd.yaml @@ -122,13 +122,6 @@ jobs: working-directory: ./data run: isort --check-only --profile black . - - name: Set up Python - uses: actions/setup-python@v5 - with: - python-version: '3.12' - cache: 'pip' - cache-dependency-path: ./data/requirements.txt - - name: Install dependencies working-directory: ./data run: | diff --git a/backend/grpcserver/server.go b/backend/grpcserver/server.go index 1f87cad..3daf56e 100644 --- a/backend/grpcserver/server.go +++ b/backend/grpcserver/server.go @@ -84,7 +84,7 @@ func (s *Server) createTask(ctx context.Context, taskID, taskType string, params } return &taskpb.TaskResponse{ TaskId: task.TaskID, - Status: taskpb.TaskStatus_TASK_STATUS_RUNNING, + Status: taskpb.TaskStatus_RUNNING, }, nil } diff --git a/backend/pkg/taskpb/task.pb.go b/backend/pkg/taskpb/task.pb.go index 4dd00c6..8e205e8 100644 --- a/backend/pkg/taskpb/task.pb.go +++ b/backend/pkg/taskpb/task.pb.go @@ -7,12 +7,11 @@ package taskpb import ( + protoreflect "google.golang.org/protobuf/reflect/protoreflect" + protoimpl "google.golang.org/protobuf/runtime/protoimpl" reflect "reflect" sync "sync" unsafe "unsafe" - - protoreflect "google.golang.org/protobuf/reflect/protoreflect" - protoimpl "google.golang.org/protobuf/runtime/protoimpl" ) const ( @@ -26,25 +25,25 @@ const ( type TaskStatus int32 const ( - TaskStatus_TASK_STATUS_RUNNING TaskStatus = 0 // 任务正在执行中 - TaskStatus_TASK_STATUS_SUCCESS TaskStatus = 1 // 任务成功完成 - TaskStatus_TASK_STATUS_FAILED TaskStatus = 2 // 任务执行失败 - TaskStatus_TASK_STATUS_CANCELED TaskStatus = 3 // 任务被取消 + TaskStatus_RUNNING TaskStatus = 0 // 任务正在执行中 + TaskStatus_SUCCESS TaskStatus = 1 // 任务成功完成 + TaskStatus_FAILED TaskStatus = 2 // 任务执行失败 + TaskStatus_CANCELLED TaskStatus = 3 // 任务被取消 ) // Enum value maps for TaskStatus. var ( TaskStatus_name = map[int32]string{ - 0: "TASK_STATUS_RUNNING", - 1: "TASK_STATUS_SUCCESS", - 2: "TASK_STATUS_FAILED", - 3: "TASK_STATUS_CANCELED", + 0: "RUNNING", + 1: "SUCCESS", + 2: "FAILED", + 3: "CANCELLED", } TaskStatus_value = map[string]int32{ - "TASK_STATUS_RUNNING": 0, - "TASK_STATUS_SUCCESS": 1, - "TASK_STATUS_FAILED": 2, - "TASK_STATUS_CANCELED": 3, + "RUNNING": 0, + "SUCCESS": 1, + "FAILED": 2, + "CANCELLED": 3, } ) @@ -135,6 +134,66 @@ func (x *CollectBinanceRequest) GetChunkSize() int32 { return 0 } +type CollectBinanceByDateRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + TaskId string `protobuf:"bytes,1,opt,name=task_id,json=taskId,proto3" json:"task_id,omitempty"` // 任务的唯一标识符 + StartTs int32 `protobuf:"varint,2,opt,name=start_ts,json=startTs,proto3" json:"start_ts,omitempty"` // 起始时间 + EndTs int32 `protobuf:"varint,3,opt,name=end_ts,json=endTs,proto3" json:"end_ts,omitempty"` // 终止时间 + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *CollectBinanceByDateRequest) Reset() { + *x = CollectBinanceByDateRequest{} + mi := &file_task_proto_msgTypes[1] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *CollectBinanceByDateRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*CollectBinanceByDateRequest) ProtoMessage() {} + +func (x *CollectBinanceByDateRequest) ProtoReflect() protoreflect.Message { + mi := &file_task_proto_msgTypes[1] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use CollectBinanceByDateRequest.ProtoReflect.Descriptor instead. +func (*CollectBinanceByDateRequest) Descriptor() ([]byte, []int) { + return file_task_proto_rawDescGZIP(), []int{1} +} + +func (x *CollectBinanceByDateRequest) GetTaskId() string { + if x != nil { + return x.TaskId + } + return "" +} + +func (x *CollectBinanceByDateRequest) GetStartTs() int32 { + if x != nil { + return x.StartTs + } + return 0 +} + +func (x *CollectBinanceByDateRequest) GetEndTs() int32 { + if x != nil { + return x.EndTs + } + return 0 +} + type CollectUniswapRequest struct { state protoimpl.MessageState `protogen:"open.v1"` TaskId string `protobuf:"bytes,1,opt,name=task_id,json=taskId,proto3" json:"task_id,omitempty"` // 任务的唯一标识符 @@ -147,7 +206,7 @@ type CollectUniswapRequest struct { func (x *CollectUniswapRequest) Reset() { *x = CollectUniswapRequest{} - mi := &file_task_proto_msgTypes[1] + mi := &file_task_proto_msgTypes[2] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -159,7 +218,7 @@ func (x *CollectUniswapRequest) String() string { func (*CollectUniswapRequest) ProtoMessage() {} func (x *CollectUniswapRequest) ProtoReflect() protoreflect.Message { - mi := &file_task_proto_msgTypes[1] + mi := &file_task_proto_msgTypes[2] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -172,7 +231,7 @@ func (x *CollectUniswapRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use CollectUniswapRequest.ProtoReflect.Descriptor instead. func (*CollectUniswapRequest) Descriptor() ([]byte, []int) { - return file_task_proto_rawDescGZIP(), []int{1} + return file_task_proto_rawDescGZIP(), []int{2} } func (x *CollectUniswapRequest) GetTaskId() string { @@ -217,7 +276,7 @@ type ProcessPricesRequest struct { func (x *ProcessPricesRequest) Reset() { *x = ProcessPricesRequest{} - mi := &file_task_proto_msgTypes[2] + mi := &file_task_proto_msgTypes[3] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -229,7 +288,7 @@ func (x *ProcessPricesRequest) String() string { func (*ProcessPricesRequest) ProtoMessage() {} func (x *ProcessPricesRequest) ProtoReflect() protoreflect.Message { - mi := &file_task_proto_msgTypes[2] + mi := &file_task_proto_msgTypes[3] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -242,7 +301,7 @@ func (x *ProcessPricesRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use ProcessPricesRequest.ProtoReflect.Descriptor instead. func (*ProcessPricesRequest) Descriptor() ([]byte, []int) { - return file_task_proto_rawDescGZIP(), []int{2} + return file_task_proto_rawDescGZIP(), []int{3} } func (x *ProcessPricesRequest) GetTaskId() string { @@ -299,7 +358,7 @@ type AnalyseRequest struct { func (x *AnalyseRequest) Reset() { *x = AnalyseRequest{} - mi := &file_task_proto_msgTypes[3] + mi := &file_task_proto_msgTypes[4] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -311,7 +370,7 @@ func (x *AnalyseRequest) String() string { func (*AnalyseRequest) ProtoMessage() {} func (x *AnalyseRequest) ProtoReflect() protoreflect.Message { - mi := &file_task_proto_msgTypes[3] + mi := &file_task_proto_msgTypes[4] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -324,7 +383,7 @@ func (x *AnalyseRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use AnalyseRequest.ProtoReflect.Descriptor instead. func (*AnalyseRequest) Descriptor() ([]byte, []int) { - return file_task_proto_rawDescGZIP(), []int{3} + return file_task_proto_rawDescGZIP(), []int{4} } func (x *AnalyseRequest) GetTaskId() string { @@ -365,7 +424,7 @@ type TaskResponse struct { func (x *TaskResponse) Reset() { *x = TaskResponse{} - mi := &file_task_proto_msgTypes[4] + mi := &file_task_proto_msgTypes[5] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -377,7 +436,7 @@ func (x *TaskResponse) String() string { func (*TaskResponse) ProtoMessage() {} func (x *TaskResponse) ProtoReflect() protoreflect.Message { - mi := &file_task_proto_msgTypes[4] + mi := &file_task_proto_msgTypes[5] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -390,7 +449,7 @@ func (x *TaskResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use TaskResponse.ProtoReflect.Descriptor instead. func (*TaskResponse) Descriptor() ([]byte, []int) { - return file_task_proto_rawDescGZIP(), []int{4} + return file_task_proto_rawDescGZIP(), []int{5} } func (x *TaskResponse) GetTaskId() string { @@ -404,7 +463,7 @@ func (x *TaskResponse) GetStatus() TaskStatus { if x != nil { return x.Status } - return TaskStatus_TASK_STATUS_RUNNING + return TaskStatus_RUNNING } var File_task_proto protoreflect.FileDescriptor @@ -417,7 +476,11 @@ const file_task_proto_rawDesc = "" + "\atask_id\x18\x01 \x01(\tR\x06taskId\x12+\n" + "\x11import_percentage\x18\x02 \x01(\x05R\x10importPercentage\x12\x1d\n" + "\n" + - "chunk_size\x18\x03 \x01(\x05R\tchunkSize\"\x85\x01\n" + + "chunk_size\x18\x03 \x01(\x05R\tchunkSize\"h\n" + + "\x1bCollectBinanceByDateRequest\x12\x17\n" + + "\atask_id\x18\x01 \x01(\tR\x06taskId\x12\x19\n" + + "\bstart_ts\x18\x02 \x01(\x05R\astartTs\x12\x15\n" + + "\x06end_ts\x18\x03 \x01(\x05R\x05endTs\"\x85\x01\n" + "\x15CollectUniswapRequest\x12\x17\n" + "\atask_id\x18\x01 \x01(\tR\x06taskId\x12!\n" + "\fpool_address\x18\x02 \x01(\tR\vpoolAddress\x12\x19\n" + @@ -441,15 +504,17 @@ const file_task_proto_rawDesc = "" + "\rstrategy_json\x18\x04 \x01(\tR\fstrategyJson\"T\n" + "\fTaskResponse\x12\x17\n" + "\atask_id\x18\x01 \x01(\tR\x06taskId\x12+\n" + - "\x06status\x18\x02 \x01(\x0e2\x13.task.v1.TaskStatusR\x06status*p\n" + + "\x06status\x18\x02 \x01(\x0e2\x13.task.v1.TaskStatusR\x06status*A\n" + + "\n" + + "TaskStatus\x12\v\n" + + "\aRUNNING\x10\x00\x12\v\n" + + "\aSUCCESS\x10\x01\x12\n" + "\n" + - "TaskStatus\x12\x17\n" + - "\x13TASK_STATUS_RUNNING\x10\x00\x12\x17\n" + - "\x13TASK_STATUS_SUCCESS\x10\x01\x12\x16\n" + - "\x12TASK_STATUS_FAILED\x10\x02\x12\x18\n" + - "\x14TASK_STATUS_CANCELED\x10\x032\xa1\x02\n" + + "\x06FAILED\x10\x02\x12\r\n" + + "\tCANCELLED\x10\x032\xf6\x02\n" + "\vTaskService\x12G\n" + - "\x0eCollectBinance\x12\x1e.task.v1.CollectBinanceRequest\x1a\x15.task.v1.TaskResponse\x12G\n" + + "\x0eCollectBinance\x12\x1e.task.v1.CollectBinanceRequest\x1a\x15.task.v1.TaskResponse\x12S\n" + + "\x14CollectBinanceByDate\x12$.task.v1.CollectBinanceByDateRequest\x1a\x15.task.v1.TaskResponse\x12G\n" + "\x0eCollectUniswap\x12\x1e.task.v1.CollectUniswapRequest\x1a\x15.task.v1.TaskResponse\x12E\n" + "\rProcessPrices\x12\x1d.task.v1.ProcessPricesRequest\x1a\x15.task.v1.TaskResponse\x129\n" + "\aAnalyse\x12\x17.task.v1.AnalyseRequest\x1a\x15.task.v1.TaskResponseB\x1bZ\x19backend/pkg/taskpb;taskpbb\x06proto3" @@ -467,29 +532,32 @@ func file_task_proto_rawDescGZIP() []byte { } var file_task_proto_enumTypes = make([]protoimpl.EnumInfo, 1) -var file_task_proto_msgTypes = make([]protoimpl.MessageInfo, 6) +var file_task_proto_msgTypes = make([]protoimpl.MessageInfo, 7) var file_task_proto_goTypes = []any{ - (TaskStatus)(0), // 0: task.v1.TaskStatus - (*CollectBinanceRequest)(nil), // 1: task.v1.CollectBinanceRequest - (*CollectUniswapRequest)(nil), // 2: task.v1.CollectUniswapRequest - (*ProcessPricesRequest)(nil), // 3: task.v1.ProcessPricesRequest - (*AnalyseRequest)(nil), // 4: task.v1.AnalyseRequest - (*TaskResponse)(nil), // 5: task.v1.TaskResponse - nil, // 6: task.v1.ProcessPricesRequest.DbOverridesEntry + (TaskStatus)(0), // 0: task.v1.TaskStatus + (*CollectBinanceRequest)(nil), // 1: task.v1.CollectBinanceRequest + (*CollectBinanceByDateRequest)(nil), // 2: task.v1.CollectBinanceByDateRequest + (*CollectUniswapRequest)(nil), // 3: task.v1.CollectUniswapRequest + (*ProcessPricesRequest)(nil), // 4: task.v1.ProcessPricesRequest + (*AnalyseRequest)(nil), // 5: task.v1.AnalyseRequest + (*TaskResponse)(nil), // 6: task.v1.TaskResponse + nil, // 7: task.v1.ProcessPricesRequest.DbOverridesEntry } var file_task_proto_depIdxs = []int32{ - 6, // 0: task.v1.ProcessPricesRequest.db_overrides:type_name -> task.v1.ProcessPricesRequest.DbOverridesEntry + 7, // 0: task.v1.ProcessPricesRequest.db_overrides:type_name -> task.v1.ProcessPricesRequest.DbOverridesEntry 0, // 1: task.v1.TaskResponse.status:type_name -> task.v1.TaskStatus 1, // 2: task.v1.TaskService.CollectBinance:input_type -> task.v1.CollectBinanceRequest - 2, // 3: task.v1.TaskService.CollectUniswap:input_type -> task.v1.CollectUniswapRequest - 3, // 4: task.v1.TaskService.ProcessPrices:input_type -> task.v1.ProcessPricesRequest - 4, // 5: task.v1.TaskService.Analyse:input_type -> task.v1.AnalyseRequest - 5, // 6: task.v1.TaskService.CollectBinance:output_type -> task.v1.TaskResponse - 5, // 7: task.v1.TaskService.CollectUniswap:output_type -> task.v1.TaskResponse - 5, // 8: task.v1.TaskService.ProcessPrices:output_type -> task.v1.TaskResponse - 5, // 9: task.v1.TaskService.Analyse:output_type -> task.v1.TaskResponse - 6, // [6:10] is the sub-list for method output_type - 2, // [2:6] is the sub-list for method input_type + 2, // 3: task.v1.TaskService.CollectBinanceByDate:input_type -> task.v1.CollectBinanceByDateRequest + 3, // 4: task.v1.TaskService.CollectUniswap:input_type -> task.v1.CollectUniswapRequest + 4, // 5: task.v1.TaskService.ProcessPrices:input_type -> task.v1.ProcessPricesRequest + 5, // 6: task.v1.TaskService.Analyse:input_type -> task.v1.AnalyseRequest + 6, // 7: task.v1.TaskService.CollectBinance:output_type -> task.v1.TaskResponse + 6, // 8: task.v1.TaskService.CollectBinanceByDate:output_type -> task.v1.TaskResponse + 6, // 9: task.v1.TaskService.CollectUniswap:output_type -> task.v1.TaskResponse + 6, // 10: task.v1.TaskService.ProcessPrices:output_type -> task.v1.TaskResponse + 6, // 11: task.v1.TaskService.Analyse:output_type -> task.v1.TaskResponse + 7, // [7:12] is the sub-list for method output_type + 2, // [2:7] is the sub-list for method input_type 2, // [2:2] is the sub-list for extension type_name 2, // [2:2] is the sub-list for extension extendee 0, // [0:2] is the sub-list for field type_name @@ -506,7 +574,7 @@ func file_task_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_task_proto_rawDesc), len(file_task_proto_rawDesc)), NumEnums: 1, - NumMessages: 6, + NumMessages: 7, NumExtensions: 0, NumServices: 1, }, diff --git a/backend/pkg/taskpb/task_grpc.pb.go b/backend/pkg/taskpb/task_grpc.pb.go index e1a3bd7..8bf42d0 100644 --- a/backend/pkg/taskpb/task_grpc.pb.go +++ b/backend/pkg/taskpb/task_grpc.pb.go @@ -1,6 +1,6 @@ // Code generated by protoc-gen-go-grpc. DO NOT EDIT. // versions: -// - protoc-gen-go-grpc v1.6.0 +// - protoc-gen-go-grpc v1.5.1 // - protoc v4.25.3 // source: task.proto @@ -8,7 +8,6 @@ package taskpb import ( context "context" - grpc "google.golang.org/grpc" codes "google.golang.org/grpc/codes" status "google.golang.org/grpc/status" @@ -20,10 +19,11 @@ import ( const _ = grpc.SupportPackageIsVersion9 const ( - TaskService_CollectBinance_FullMethodName = "/task.v1.TaskService/CollectBinance" - TaskService_CollectUniswap_FullMethodName = "/task.v1.TaskService/CollectUniswap" - TaskService_ProcessPrices_FullMethodName = "/task.v1.TaskService/ProcessPrices" - TaskService_Analyse_FullMethodName = "/task.v1.TaskService/Analyse" + TaskService_CollectBinance_FullMethodName = "/task.v1.TaskService/CollectBinance" + TaskService_CollectBinanceByDate_FullMethodName = "/task.v1.TaskService/CollectBinanceByDate" + TaskService_CollectUniswap_FullMethodName = "/task.v1.TaskService/CollectUniswap" + TaskService_ProcessPrices_FullMethodName = "/task.v1.TaskService/ProcessPrices" + TaskService_Analyse_FullMethodName = "/task.v1.TaskService/Analyse" ) // TaskServiceClient is the client API for TaskService service. @@ -34,6 +34,8 @@ const ( type TaskServiceClient interface { // 收集币安数据 CollectBinance(ctx context.Context, in *CollectBinanceRequest, opts ...grpc.CallOption) (*TaskResponse, error) + // 按日期收集币安数据 + CollectBinanceByDate(ctx context.Context, in *CollectBinanceByDateRequest, opts ...grpc.CallOption) (*TaskResponse, error) // 收集Uniswap数据 CollectUniswap(ctx context.Context, in *CollectUniswapRequest, opts ...grpc.CallOption) (*TaskResponse, error) // 处理价格数据 @@ -60,6 +62,16 @@ func (c *taskServiceClient) CollectBinance(ctx context.Context, in *CollectBinan return out, nil } +func (c *taskServiceClient) CollectBinanceByDate(ctx context.Context, in *CollectBinanceByDateRequest, opts ...grpc.CallOption) (*TaskResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(TaskResponse) + err := c.cc.Invoke(ctx, TaskService_CollectBinanceByDate_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + func (c *taskServiceClient) CollectUniswap(ctx context.Context, in *CollectUniswapRequest, opts ...grpc.CallOption) (*TaskResponse, error) { cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) out := new(TaskResponse) @@ -98,6 +110,8 @@ func (c *taskServiceClient) Analyse(ctx context.Context, in *AnalyseRequest, opt type TaskServiceServer interface { // 收集币安数据 CollectBinance(context.Context, *CollectBinanceRequest) (*TaskResponse, error) + // 按日期收集币安数据 + CollectBinanceByDate(context.Context, *CollectBinanceByDateRequest) (*TaskResponse, error) // 收集Uniswap数据 CollectUniswap(context.Context, *CollectUniswapRequest) (*TaskResponse, error) // 处理价格数据 @@ -115,16 +129,19 @@ type TaskServiceServer interface { type UnimplementedTaskServiceServer struct{} func (UnimplementedTaskServiceServer) CollectBinance(context.Context, *CollectBinanceRequest) (*TaskResponse, error) { - return nil, status.Error(codes.Unimplemented, "method CollectBinance not implemented") + return nil, status.Errorf(codes.Unimplemented, "method CollectBinance not implemented") +} +func (UnimplementedTaskServiceServer) CollectBinanceByDate(context.Context, *CollectBinanceByDateRequest) (*TaskResponse, error) { + return nil, status.Errorf(codes.Unimplemented, "method CollectBinanceByDate not implemented") } func (UnimplementedTaskServiceServer) CollectUniswap(context.Context, *CollectUniswapRequest) (*TaskResponse, error) { - return nil, status.Error(codes.Unimplemented, "method CollectUniswap not implemented") + return nil, status.Errorf(codes.Unimplemented, "method CollectUniswap not implemented") } func (UnimplementedTaskServiceServer) ProcessPrices(context.Context, *ProcessPricesRequest) (*TaskResponse, error) { - return nil, status.Error(codes.Unimplemented, "method ProcessPrices not implemented") + return nil, status.Errorf(codes.Unimplemented, "method ProcessPrices not implemented") } func (UnimplementedTaskServiceServer) Analyse(context.Context, *AnalyseRequest) (*TaskResponse, error) { - return nil, status.Error(codes.Unimplemented, "method Analyse not implemented") + return nil, status.Errorf(codes.Unimplemented, "method Analyse not implemented") } func (UnimplementedTaskServiceServer) mustEmbedUnimplementedTaskServiceServer() {} func (UnimplementedTaskServiceServer) testEmbeddedByValue() {} @@ -137,7 +154,7 @@ type UnsafeTaskServiceServer interface { } func RegisterTaskServiceServer(s grpc.ServiceRegistrar, srv TaskServiceServer) { - // If the following call panics, it indicates UnimplementedTaskServiceServer was + // If the following call pancis, it indicates UnimplementedTaskServiceServer was // embedded by pointer and is nil. This will cause panics if an // unimplemented method is ever invoked, so we test this at initialization // time to prevent it from happening at runtime later due to I/O. @@ -165,6 +182,24 @@ func _TaskService_CollectBinance_Handler(srv interface{}, ctx context.Context, d return interceptor(ctx, in, info, handler) } +func _TaskService_CollectBinanceByDate_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(CollectBinanceByDateRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(TaskServiceServer).CollectBinanceByDate(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: TaskService_CollectBinanceByDate_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(TaskServiceServer).CollectBinanceByDate(ctx, req.(*CollectBinanceByDateRequest)) + } + return interceptor(ctx, in, info, handler) +} + func _TaskService_CollectUniswap_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { in := new(CollectUniswapRequest) if err := dec(in); err != nil { @@ -230,6 +265,10 @@ var TaskService_ServiceDesc = grpc.ServiceDesc{ MethodName: "CollectBinance", Handler: _TaskService_CollectBinance_Handler, }, + { + MethodName: "CollectBinanceByDate", + Handler: _TaskService_CollectBinanceByDate_Handler, + }, { MethodName: "CollectUniswap", Handler: _TaskService_CollectUniswap_Handler, diff --git a/data/.gitignore b/data/.gitignore index 1fe52a1..1a01e29 100644 --- a/data/.gitignore +++ b/data/.gitignore @@ -218,4 +218,5 @@ __marimo__/ # files logs/ *.csv -ETHUSDT-trades-2025-09.zip \ No newline at end of file +ETHUSDT-trades-2025-09.zip +allure-results/ \ No newline at end of file diff --git a/data/block_chain/analyse.py b/data/block_chain/analyse.py index 557bcd4..2516d7d 100644 --- a/data/block_chain/analyse.py +++ b/data/block_chain/analyse.py @@ -367,7 +367,11 @@ def run_analyse(task_id: Optional[str] = None, config_json: Optional[str] = None raise else: logger.info(f"分析完成,发现 {len(opportunities)} 条机会") - update_task_status(task_id, 1) + # 在标记成功前,再次检查任务是否被取消 + if check_task(task_id): + logger.info(f"任务 {task_id} 已取消,不标记为成功") + else: + update_task_status(task_id, 1) conn.close() diff --git a/data/block_chain/collect_binance.py b/data/block_chain/collect_binance.py index ce66cfe..dea82d0 100644 --- a/data/block_chain/collect_binance.py +++ b/data/block_chain/collect_binance.py @@ -1,4 +1,3 @@ -import argparse import io import math import os @@ -8,7 +7,7 @@ import traceback import zipfile from datetime import datetime, timedelta, timezone -from typing import Optional +from typing import Any, Optional import pandas as pd import psycopg2 @@ -74,10 +73,11 @@ def process_chunk( chunk_index: int, rows_counter, target_rows: Optional[int], + conn: Optional[Any] = None, ): """ 描述:处理单个分块:预处理数据并写入数据库 - 参数:task_id: 任务ID, chunk_data: 分块数据, chunk_index: 分块索引, rows_counter: 计数器, target_rows: 目标行数 + 参数:task_id: 任务ID, chunk_data: 分块数据, chunk_index: 分块索引, rows_counter: 计数器, target_rows: 目标行数, conn: 数据库连接(可选) 返回值:成功标志, 处理行数, 导入行数, 是否停止标志 """ original_chunk_len = len(chunk_data) @@ -105,16 +105,23 @@ def process_chunk( csv_buffer.seek(0) columns = "id, price, qty, quote_qty, trade_time, is_buyer_maker, is_best_match" copy_sql = f"COPY binance_trades ({columns}) FROM STDIN WITH (FORMAT CSV)" - with psycopg2.connect( - host=db_config["host"], - port=db_config["port"], - dbname=db_config["database"], - user=db_config["username"], - password=db_config["password"], - ) as conn: + + # 如果提供了连接,使用它;否则创建新连接 + if conn is not None: with conn.cursor() as cursor: cursor.copy_expert(sql=copy_sql, file=csv_buffer) - conn.commit() + # 不在这里提交,由调用者控制事务 + else: + with psycopg2.connect( + host=db_config["host"], + port=db_config["port"], + dbname=db_config["database"], + user=db_config["username"], + password=db_config["password"], + ) as new_conn: + with new_conn.cursor() as cursor: + cursor.copy_expert(sql=copy_sql, file=csv_buffer) + new_conn.commit() rows_imported = len(chunk) rows_counter[0] += original_chunk_len rows_counter[1] += rows_imported @@ -132,10 +139,11 @@ def import_data_to_database( target_rows: Optional[int], total_lines: Optional[int], chunk_size: int, + conn: Optional[Any] = None, ): """ 描述:主导入逻辑:读取CSV,分块处理并写入数据库。 - 参数:target_rows: 目标行数, total_lines: 总行数, chunk_size: 分块大小 + 参数:target_rows: 目标行数, total_lines: 总行数, chunk_size: 分块大小, conn: 数据库连接(可选) 返回值:处理行数, 导入行数 """ rows_counter = [0, 0] @@ -164,6 +172,7 @@ def import_data_to_database( i, rows_counter, target_rows, + conn, ) if stop_flag and not should_stop: logger.info( @@ -196,27 +205,75 @@ def _calc_target_rows( def collect_binance( task_id: str, csv_path: str, import_percentage: int, chunk_size: int ): + """ + 描述:收集 Binance 数据(作为事务处理,如果任务取消则完全回滚) + 参数: + task_id: 任务ID + csv_path: CSV文件路径 + import_percentage: 导入百分比 + chunk_size: 分块大小 + 返回值:导入的总行数 + """ + conn = None try: start_time = time.time() total_lines = count_lines(task_id, csv_path) if check_task(task_id): logger.info(f"任务 {task_id} 已取消,停止导入 Binance 数据") return 0 + + # 创建数据库连接并开始事务 + conn = psycopg2.connect( + host=db_config["host"], + port=db_config["port"], + dbname=db_config["database"], + user=db_config["username"], + password=db_config["password"], + ) + conn.autocommit = False # 禁用自动提交,使用事务 + logger.info("已开启数据库事务,所有导入操作将在事务中执行") + target_rows = _calc_target_rows(total_lines, import_percentage) rows_counter = import_data_to_database( - task_id, csv_path, target_rows, total_lines, chunk_size + task_id, csv_path, target_rows, total_lines, chunk_size, conn ) total_time = time.time() - start_time + if check_task(task_id): - logger.info(f"任务 {task_id} 已取消,停止导入 Binance 数据") + logger.info(f"任务 {task_id} 已取消,回滚所有已导入的数据") + conn.rollback() + logger.info("已回滚所有数据") return 0 + + # 所有数据导入成功,提交事务 + conn.commit() + logger.info("事务已提交,所有数据已成功导入") logger.info(f"成功导入 {rows_counter[1]} 行,耗时 {total_time:.2f}s") - update_task_status(task_id, "SUCCESS") + # 在标记成功前,再次检查任务是否被取消 + if check_task(task_id): + logger.info(f"任务 {task_id} 已取消,不标记为成功") + else: + update_task_status(task_id, "SUCCESS") return rows_counter[1] except Exception as e: logger.error(f"导入 Binance 数据失败: {e}") + traceback.print_exc(file=sys.stderr) + # 确保在异常情况下回滚事务 + if conn is not None: + try: + conn.rollback() + logger.info("发生异常,已回滚所有数据") + except Exception as rollback_error: + logger.error(f"回滚事务失败: {rollback_error}") update_task_status(task_id, "FAILED") raise + finally: + # 确保关闭数据库连接 + if conn is not None: + try: + conn.close() + except Exception as close_error: + logger.warning(f"关闭数据库连接失败: {close_error}") def download_binance_file( @@ -290,7 +347,7 @@ def collect_binance_by_date( chunk_size: int = 1000000, ) -> int: """ - 描述:按日期范围收集币安数据 + 描述:按日期范围收集币安数据(作为事务处理,如果任务取消则完全回滚) 参数: task_id: 任务ID start_ts: 起始时间戳(秒级) @@ -299,6 +356,7 @@ def collect_binance_by_date( chunk_size: 分块大小,默认1000000 返回值:导入的总行数 """ + conn = None try: start_time = time.time() @@ -308,6 +366,17 @@ def collect_binance_by_date( logger.info(f"开始按日期收集币安数据: {start_date.date()} 到 {end_date.date()}") + # 创建数据库连接并开始事务 + conn = psycopg2.connect( + host=db_config["host"], + port=db_config["port"], + dbname=db_config["database"], + user=db_config["username"], + password=db_config["password"], + ) + conn.autocommit = False # 禁用自动提交,使用事务 + logger.info("已开启数据库事务,所有导入操作将在事务中执行") + total_rows_imported = 0 temp_files = [] # 记录临时文件,用于清理 @@ -317,8 +386,10 @@ def collect_binance_by_date( while current_date <= end_date_only: if check_task(task_id): - logger.info(f"任务 {task_id} 已取消,停止收集 Binance 数据") - break + logger.info(f"任务 {task_id} 已取消,回滚所有已导入的数据") + conn.rollback() + logger.info("已回滚所有数据") + return 0 date_str = current_date.strftime("%Y-%m-%d") logger.info(f"正在处理日期: {date_str}") @@ -336,14 +407,18 @@ def collect_binance_by_date( try: # 导入数据(导入全部数据,不限制百分比) + # 使用共享的数据库连接,所有操作在同一事务中 rows_counter = import_data_to_database( - task_id, csv_path, None, None, chunk_size + task_id, csv_path, None, None, chunk_size, conn ) total_rows_imported += rows_counter[1] logger.info(f"日期 {date_str} 导入完成,导入 {rows_counter[1]} 行") except Exception as e: logger.error(f"导入日期 {date_str} 的数据失败: {e}") - # 继续处理下一个日期,不中断整个任务 + # 发生错误,回滚事务 + conn.rollback() + logger.error("已回滚所有数据") + raise # 清理临时文件 try: @@ -360,18 +435,42 @@ def collect_binance_by_date( total_time = time.time() - start_time if check_task(task_id): - logger.info(f"任务 {task_id} 已取消,停止收集 Binance 数据") + logger.info(f"任务 {task_id} 已取消,回滚所有已导入的数据") + conn.rollback() + logger.info("已回滚所有数据") return 0 + # 所有数据导入成功,提交事务 + conn.commit() + logger.info("事务已提交,所有数据已成功导入") + logger.info(f"成功导入 {total_rows_imported} 行,耗时 {total_time:.2f}s") - update_task_status(task_id, "SUCCESS") + # 在标记成功前,再次检查任务是否被取消 + if check_task(task_id): + logger.info(f"任务 {task_id} 已取消,不标记为成功") + else: + update_task_status(task_id, "SUCCESS") return total_rows_imported except Exception as e: logger.error(f"按日期收集 Binance 数据失败: {e}") traceback.print_exc(file=sys.stderr) + # 确保在异常情况下回滚事务 + if conn is not None: + try: + conn.rollback() + logger.info("发生异常,已回滚所有数据") + except Exception as rollback_error: + logger.error(f"回滚事务失败: {rollback_error}") update_task_status(task_id, "FAILED") raise + finally: + # 确保关闭数据库连接 + if conn is not None: + try: + conn.close() + except Exception as close_error: + logger.warning(f"关闭数据库连接失败: {close_error}") if __name__ == "__main__": diff --git a/data/block_chain/collect_uniswap.py b/data/block_chain/collect_uniswap.py index e940ff2..e7b8ca7 100644 --- a/data/block_chain/collect_uniswap.py +++ b/data/block_chain/collect_uniswap.py @@ -1,5 +1,5 @@ import time -from typing import Any, Iterable +from typing import Any, Iterable, Optional import pandas as pd import psycopg2 @@ -89,13 +89,14 @@ def fetch_all_swaps(task_id: str, pool_address: str, start_ts: int, end_ts: int) def process_and_store_uniswap_data( - task_id: str, swaps_data: Iterable[dict[str, Any]] + task_id: str, swaps_data: Iterable[dict[str, Any]], conn: Optional[Any] = None ) -> int: """ 描述:处理数据并存入数据库。 参数: task_id: 任务ID swaps_data: Uniswap数据 + conn: 数据库连接(可选),如果提供则使用该连接,否则创建新连接 返回值:写入的记录数量 """ swaps = list(swaps_data) @@ -128,35 +129,97 @@ def process_and_store_uniswap_data( logger.info("没有可写入的数据。") return 0 - with psycopg2.connect( - host=db_config["host"], - port=db_config["port"], - dbname=db_config["database"], - user=db_config["username"], - password=db_config["password"], - ) as conn, conn.cursor() as cur: - insert_sql = """ - INSERT INTO uniswap_swaps (block_time, price, amount_eth, amount_usdt, gas_price, tx_hash) - VALUES %s - """ - execute_values(cur, insert_sql, records, page_size=1000) + # 如果提供了连接,使用它;否则创建新连接 + if conn is not None: + with conn.cursor() as cur: + insert_sql = """ + INSERT INTO uniswap_swaps (block_time, price, amount_eth, amount_usdt, gas_price, tx_hash) + VALUES %s + """ + execute_values(cur, insert_sql, records, page_size=1000) + # 不在这里提交,由调用者控制事务 + else: + with psycopg2.connect( + host=db_config["host"], + port=db_config["port"], + dbname=db_config["database"], + user=db_config["username"], + password=db_config["password"], + ) as new_conn, new_conn.cursor() as cur: + insert_sql = """ + INSERT INTO uniswap_swaps (block_time, price, amount_eth, amount_usdt, gas_price, tx_hash) + VALUES %s + """ + execute_values(cur, insert_sql, records, page_size=1000) logger.info(f"成功写入 {len(records)} 条 Uniswap 记录。") return len(records) def collect_uniswap(task_id: str, pool_address: str, start_ts: int, end_ts: int) -> int: + """ + 描述:收集 Uniswap 数据(作为事务处理,如果任务取消则完全回滚) + 参数: + task_id: 任务ID + pool_address: 池地址 + start_ts: 起始时间戳(秒级) + end_ts: 终止时间戳(秒级) + 返回值:导入的总行数 + """ + conn = None try: + # 创建数据库连接并开始事务 + conn = psycopg2.connect( + host=db_config["host"], + port=db_config["port"], + dbname=db_config["database"], + user=db_config["username"], + password=db_config["password"], + ) + conn.autocommit = False # 禁用自动提交,使用事务 + logger.info("已开启数据库事务,所有导入操作将在事务中执行") + swaps = fetch_all_swaps(task_id, pool_address, start_ts, end_ts) if check_task(task_id): - logger.info(f"任务 {task_id} 已取消,停止写入 Uniswap 数据") + logger.info(f"任务 {task_id} 已取消,回滚所有已导入的数据") + conn.rollback() + logger.info("已回滚所有数据") return 0 - rows_counter = process_and_store_uniswap_data(task_id, swaps) - update_task_status(task_id, "SUCCESS") + + rows_counter = process_and_store_uniswap_data(task_id, swaps, conn) + + if check_task(task_id): + logger.info(f"任务 {task_id} 已取消,回滚所有已导入的数据") + conn.rollback() + logger.info("已回滚所有数据") + return 0 + + # 所有数据导入成功,提交事务 + conn.commit() + logger.info("事务已提交,所有数据已成功导入") + # 在标记成功前,再次检查任务是否被取消 + if check_task(task_id): + logger.info(f"任务 {task_id} 已取消,不标记为成功") + else: + update_task_status(task_id, "SUCCESS") return rows_counter except Exception as e: logger.error(f"获取Uniswap数据失败: {e}") + # 确保在异常情况下回滚事务 + if conn is not None: + try: + conn.rollback() + logger.info("发生异常,已回滚所有数据") + except Exception as rollback_error: + logger.error(f"回滚事务失败: {rollback_error}") update_task_status(task_id, "FAILED") return 0 + finally: + # 确保关闭数据库连接 + if conn is not None: + try: + conn.close() + except Exception as close_error: + logger.warning(f"关闭数据库连接失败: {close_error}") if __name__ == "__main__": diff --git a/data/block_chain/process_prices.py b/data/block_chain/process_prices.py index eda696f..1dd763a 100644 --- a/data/block_chain/process_prices.py +++ b/data/block_chain/process_prices.py @@ -171,8 +171,14 @@ def run_process_prices(task_id: str, **kwargs: Any): update_task_status(task_id, 2) raise else: - logger.info(f"聚合完成,共写入 {len(df_final)} 条记录,耗时 {duration:.2f}s") - update_task_status(task_id, 1) + # 在标记成功前,再次检查任务是否被取消 + if check_task(task_id): + logger.info(f"任务 {task_id} 已取消,不标记为成功") + else: + logger.info( + f"聚合完成,共写入 {len(df_final)} 条记录,耗时 {duration:.2f}s" + ) + update_task_status(task_id, 1) finally: conn.close() diff --git a/data/block_chain/task.py b/data/block_chain/task.py index e18d4e6..71cf741 100644 --- a/data/block_chain/task.py +++ b/data/block_chain/task.py @@ -51,6 +51,7 @@ def update_task_status(task_id: str, status: str): "UPDATE tasks SET status = %s WHERE task_id = %s", (status, str(task_id)), ) + conn.commit() except Exception as e: logger.error(f"更新任务 {task_id} 状态失败: {e}") diff --git a/data/config/config.yaml b/data/config/config.yaml index 9a1dc2d..164baf0 100644 --- a/data/config/config.yaml +++ b/data/config/config.yaml @@ -9,3 +9,9 @@ the_graph: api_key: 9f9faba5da813868926b3337fb728af5 graph_api_url: https://gateway.thegraph.com/api/subgraphs/id/5zvR82QoaXYFyDEKLZ9t6v9adgnptxYpKpSbxtgVENFV uniswap_pool_address: 0x11b815efb8f581194ae79006d24e0d814b7697f6 + +rabbitmq: + host: localhost + port: 5672 + username: admin + password: 123456 \ No newline at end of file diff --git a/data/protos/task_pb2.py b/data/protos/task_pb2.py index 3eb0009..35bb941 100644 --- a/data/protos/task_pb2.py +++ b/data/protos/task_pb2.py @@ -19,7 +19,7 @@ DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( - b'\n\x11protos/task.proto\x12\x07task.v1"W\n\x15\x43ollectBinanceRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x19\n\x11import_percentage\x18\x02 \x01(\x05\x12\x12\n\nchunk_size\x18\x03 \x01(\x05"O\n\x1a\x43ollectBinaceByDateRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x10\n\x08start_ts\x18\x02 \x01(\x05\x12\x0e\n\x06\x65nd_ts\x18\x03 \x01(\x05"`\n\x15\x43ollectUniswapRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x14\n\x0cpool_address\x18\x02 \x01(\t\x12\x10\n\x08start_ts\x18\x03 \x01(\x05\x12\x0e\n\x06\x65nd_ts\x18\x04 \x01(\x05"\xf8\x01\n\x14ProcessPricesRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x12\n\nstart_date\x18\x02 \x01(\x05\x12\x10\n\x08\x65nd_date\x18\x03 \x01(\x05\x12\x1c\n\x14\x61ggregation_interval\x18\x04 \x01(\t\x12\x11\n\toverwrite\x18\x05 \x01(\x08\x12\x44\n\x0c\x64\x62_overrides\x18\x06 \x03(\x0b\x32..task.v1.ProcessPricesRequest.DbOverridesEntry\x1a\x32\n\x10\x44\x62OverridesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\t:\x02\x38\x01"]\n\x0e\x41nalyseRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x10\n\x08\x62\x61tch_id\x18\x02 \x01(\x05\x12\x11\n\toverwrite\x18\x03 \x01(\x08\x12\x15\n\rstrategy_json\x18\x04 \x01(\t"D\n\x0cTaskResponse\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12#\n\x06status\x18\x02 \x01(\x0e\x32\x13.task.v1.TaskStatus*p\n\nTaskStatus\x12\x17\n\x13TASK_STATUS_RUNNING\x10\x00\x12\x17\n\x13TASK_STATUS_SUCCESS\x10\x01\x12\x16\n\x12TASK_STATUS_FAILED\x10\x02\x12\x18\n\x14TASK_STATUS_CANCELED\x10\x03\x32\xf5\x02\n\x0bTaskService\x12G\n\x0e\x43ollectBinance\x12\x1e.task.v1.CollectBinanceRequest\x1a\x15.task.v1.TaskResponse\x12R\n\x14\x43ollectBinanceByDate\x12#.task.v1.CollectBinaceByDateRequest\x1a\x15.task.v1.TaskResponse\x12G\n\x0e\x43ollectUniswap\x12\x1e.task.v1.CollectUniswapRequest\x1a\x15.task.v1.TaskResponse\x12\x45\n\rProcessPrices\x12\x1d.task.v1.ProcessPricesRequest\x1a\x15.task.v1.TaskResponse\x12\x39\n\x07\x41nalyse\x12\x17.task.v1.AnalyseRequest\x1a\x15.task.v1.TaskResponseB\x1bZ\x19\x62\x61\x63kend/pkg/taskpb;taskpbb\x06proto3' + b'\n\x11protos/task.proto\x12\x07task.v1"W\n\x15\x43ollectBinanceRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x19\n\x11import_percentage\x18\x02 \x01(\x05\x12\x12\n\nchunk_size\x18\x03 \x01(\x05"P\n\x1b\x43ollectBinanceByDateRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x10\n\x08start_ts\x18\x02 \x01(\x05\x12\x0e\n\x06\x65nd_ts\x18\x03 \x01(\x05"`\n\x15\x43ollectUniswapRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x14\n\x0cpool_address\x18\x02 \x01(\t\x12\x10\n\x08start_ts\x18\x03 \x01(\x05\x12\x0e\n\x06\x65nd_ts\x18\x04 \x01(\x05"\xf8\x01\n\x14ProcessPricesRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x12\n\nstart_date\x18\x02 \x01(\x05\x12\x10\n\x08\x65nd_date\x18\x03 \x01(\x05\x12\x1c\n\x14\x61ggregation_interval\x18\x04 \x01(\t\x12\x11\n\toverwrite\x18\x05 \x01(\x08\x12\x44\n\x0c\x64\x62_overrides\x18\x06 \x03(\x0b\x32..task.v1.ProcessPricesRequest.DbOverridesEntry\x1a\x32\n\x10\x44\x62OverridesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\t:\x02\x38\x01"]\n\x0e\x41nalyseRequest\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12\x10\n\x08\x62\x61tch_id\x18\x02 \x01(\x05\x12\x11\n\toverwrite\x18\x03 \x01(\x08\x12\x15\n\rstrategy_json\x18\x04 \x01(\t"D\n\x0cTaskResponse\x12\x0f\n\x07task_id\x18\x01 \x01(\t\x12#\n\x06status\x18\x02 \x01(\x0e\x32\x13.task.v1.TaskStatus*K\n\nTaskStatus\x12\x08\n\x04WAIT\x10\x00\x12\x0b\n\x07RUNNING\x10\x01\x12\x0b\n\x07SUCCESS\x10\x02\x12\n\n\x06\x46\x41ILED\x10\x03\x12\r\n\tCANCELLED\x10\x04\x32\xf6\x02\n\x0bTaskService\x12G\n\x0e\x43ollectBinance\x12\x1e.task.v1.CollectBinanceRequest\x1a\x15.task.v1.TaskResponse\x12S\n\x14\x43ollectBinanceByDate\x12$.task.v1.CollectBinanceByDateRequest\x1a\x15.task.v1.TaskResponse\x12G\n\x0e\x43ollectUniswap\x12\x1e.task.v1.CollectUniswapRequest\x1a\x15.task.v1.TaskResponse\x12\x45\n\rProcessPrices\x12\x1d.task.v1.ProcessPricesRequest\x1a\x15.task.v1.TaskResponse\x12\x39\n\x07\x41nalyse\x12\x17.task.v1.AnalyseRequest\x1a\x15.task.v1.TaskResponseB\x1bZ\x19\x62\x61\x63kend/pkg/taskpb;taskpbb\x06proto3' ) _globals = globals() @@ -30,22 +30,22 @@ _globals["DESCRIPTOR"]._serialized_options = b"Z\031backend/pkg/taskpb;taskpb" _globals["_PROCESSPRICESREQUEST_DBOVERRIDESENTRY"]._loaded_options = None _globals["_PROCESSPRICESREQUEST_DBOVERRIDESENTRY"]._serialized_options = b"8\001" - _globals["_TASKSTATUS"]._serialized_start = 714 - _globals["_TASKSTATUS"]._serialized_end = 826 + _globals["_TASKSTATUS"]._serialized_start = 715 + _globals["_TASKSTATUS"]._serialized_end = 790 _globals["_COLLECTBINANCEREQUEST"]._serialized_start = 30 _globals["_COLLECTBINANCEREQUEST"]._serialized_end = 117 - _globals["_COLLECTBINACEBYDATEREQUEST"]._serialized_start = 119 - _globals["_COLLECTBINACEBYDATEREQUEST"]._serialized_end = 198 - _globals["_COLLECTUNISWAPREQUEST"]._serialized_start = 200 - _globals["_COLLECTUNISWAPREQUEST"]._serialized_end = 296 - _globals["_PROCESSPRICESREQUEST"]._serialized_start = 299 - _globals["_PROCESSPRICESREQUEST"]._serialized_end = 547 - _globals["_PROCESSPRICESREQUEST_DBOVERRIDESENTRY"]._serialized_start = 497 - _globals["_PROCESSPRICESREQUEST_DBOVERRIDESENTRY"]._serialized_end = 547 - _globals["_ANALYSEREQUEST"]._serialized_start = 549 - _globals["_ANALYSEREQUEST"]._serialized_end = 642 - _globals["_TASKRESPONSE"]._serialized_start = 644 - _globals["_TASKRESPONSE"]._serialized_end = 712 - _globals["_TASKSERVICE"]._serialized_start = 829 - _globals["_TASKSERVICE"]._serialized_end = 1202 + _globals["_COLLECTBINANCEBYDATEREQUEST"]._serialized_start = 119 + _globals["_COLLECTBINANCEBYDATEREQUEST"]._serialized_end = 199 + _globals["_COLLECTUNISWAPREQUEST"]._serialized_start = 201 + _globals["_COLLECTUNISWAPREQUEST"]._serialized_end = 297 + _globals["_PROCESSPRICESREQUEST"]._serialized_start = 300 + _globals["_PROCESSPRICESREQUEST"]._serialized_end = 548 + _globals["_PROCESSPRICESREQUEST_DBOVERRIDESENTRY"]._serialized_start = 498 + _globals["_PROCESSPRICESREQUEST_DBOVERRIDESENTRY"]._serialized_end = 548 + _globals["_ANALYSEREQUEST"]._serialized_start = 550 + _globals["_ANALYSEREQUEST"]._serialized_end = 643 + _globals["_TASKRESPONSE"]._serialized_start = 645 + _globals["_TASKRESPONSE"]._serialized_end = 713 + _globals["_TASKSERVICE"]._serialized_start = 793 + _globals["_TASKSERVICE"]._serialized_end = 1167 # @@protoc_insertion_point(module_scope) diff --git a/data/protos/task_pb2_grpc.py b/data/protos/task_pb2_grpc.py index b32ca65..d8885ad 100644 --- a/data/protos/task_pb2_grpc.py +++ b/data/protos/task_pb2_grpc.py @@ -46,7 +46,7 @@ def __init__(self, channel): ) self.CollectBinanceByDate = channel.unary_unary( "/task.v1.TaskService/CollectBinanceByDate", - request_serializer=protos_dot_task__pb2.CollectBinaceByDateRequest.SerializeToString, + request_serializer=protos_dot_task__pb2.CollectBinanceByDateRequest.SerializeToString, response_deserializer=protos_dot_task__pb2.TaskResponse.FromString, _registered_method=True, ) @@ -113,7 +113,7 @@ def add_TaskServiceServicer_to_server(servicer, server): ), "CollectBinanceByDate": grpc.unary_unary_rpc_method_handler( servicer.CollectBinanceByDate, - request_deserializer=protos_dot_task__pb2.CollectBinaceByDateRequest.FromString, + request_deserializer=protos_dot_task__pb2.CollectBinanceByDateRequest.FromString, response_serializer=protos_dot_task__pb2.TaskResponse.SerializeToString, ), "CollectUniswap": grpc.unary_unary_rpc_method_handler( @@ -190,7 +190,7 @@ def CollectBinanceByDate( request, target, "/task.v1.TaskService/CollectBinanceByDate", - protos_dot_task__pb2.CollectBinaceByDateRequest.SerializeToString, + protos_dot_task__pb2.CollectBinanceByDateRequest.SerializeToString, protos_dot_task__pb2.TaskResponse.FromString, options, channel_credentials, diff --git a/data/requirements.txt b/data/requirements.txt index 60e091f..010940f 100644 --- a/data/requirements.txt +++ b/data/requirements.txt @@ -2,6 +2,7 @@ grpcio~=1.76.0 loguru~=0.7.3 numpy~=2.3.5 pandas~=2.3.3 +pika~=1.3.2 protobuf~=6.33.1 psycopg2~=2.9.11 psycopg2_binary~=2.9.11 diff --git a/data/server.py b/data/server.py index 86a2660..938ec35 100644 --- a/data/server.py +++ b/data/server.py @@ -3,14 +3,18 @@ import os import sys import threading +import time from concurrent import futures import grpc +import pika +import pika.exceptions import psycopg2 import yaml from loguru import logger from block_chain import analyse, collect_binance, collect_uniswap, process_prices +from block_chain.task import check_task # 导入生成的代码 from protos.task_pb2 import TaskResponse, TaskStatus @@ -24,13 +28,18 @@ ) as cfg_file: CONFIG = yaml.safe_load(cfg_file) DB_CONFIG = CONFIG.get("db", {}) +RABBITMQ_CONFIG = CONFIG.get("rabbitmq", {}) WORKER_PORT = str(CONFIG.get("worker_port", 50052)) +# 任务队列名称 +TASK_QUEUE_NAME = "task_queue" + STATUS_LABELS = { - TaskStatus.TASK_STATUS_RUNNING: "RUNNING", - TaskStatus.TASK_STATUS_SUCCESS: "SUCCESS", - TaskStatus.TASK_STATUS_FAILED: "FAILED", - TaskStatus.TASK_STATUS_CANCELED: "CANCELED", + TaskStatus.WAIT: "WAIT", + TaskStatus.RUNNING: "RUNNING", + TaskStatus.SUCCESS: "SUCCESS", + TaskStatus.FAILED: "FAILED", + TaskStatus.CANCELLED: "CANCELLED", } @@ -63,17 +72,37 @@ def log_task_event(task_id: str, level: str, message: str): def mark_task_started(task_id: str): + """将任务状态更新为 RUNNING 并设置开始时间""" if not task_id: return try: with _get_db_connection() as conn, conn.cursor() as cur: cur.execute( - "UPDATE tasks SET started_at = COALESCE(started_at, NOW()) WHERE task_id = %s", - (task_id,), + """ + UPDATE tasks + SET status = %s, started_at = COALESCE(started_at, NOW()) + WHERE task_id = %s + """, + (STATUS_LABELS.get(TaskStatus.RUNNING), task_id), ) conn.commit() except Exception as exc: - logger.warning("更新任务开始时间失败: %s", exc) + logger.warning("更新任务开始时间和状态失败: %s", exc) + + +def mark_task_waiting(task_id: str): + """将任务状态设置为 WAIT""" + if not task_id: + return + try: + with _get_db_connection() as conn, conn.cursor() as cur: + cur.execute( + "UPDATE tasks SET status = %s WHERE task_id = %s", + (STATUS_LABELS.get(TaskStatus.WAIT), task_id), + ) + conn.commit() + except Exception as exc: + logger.warning("更新任务状态为 WAIT 失败: %s", exc) def mark_task_finished(task_id: str, status: TaskStatus, summary: str | None = None): @@ -100,13 +129,298 @@ def mark_task_finished(task_id: str, status: TaskStatus, summary: str | None = N logger.warning("更新任务结束状态失败: %s", exc) +class RabbitMQManager: + """RabbitMQ 连接和队列管理器(线程安全)""" + + def __init__(self): + self._publish_connection = None + self._publish_channel = None + self._publish_lock = threading.Lock() + self._consume_connection = None + self._consume_channel = None + + def _get_connection_parameters(self): + """获取连接参数""" + username = str(RABBITMQ_CONFIG.get("username", "guest")) + password = str(RABBITMQ_CONFIG.get("password", "guest")) + credentials = pika.PlainCredentials(username, password) + return pika.ConnectionParameters( + host=RABBITMQ_CONFIG.get("host", "localhost"), + port=int(RABBITMQ_CONFIG.get("port", 5672)), + credentials=credentials, + heartbeat=600, + blocked_connection_timeout=300, + ) + + def connect(self): + """连接到 RabbitMQ 服务器(初始化发布连接)""" + try: + parameters = self._get_connection_parameters() + self._publish_connection = pika.BlockingConnection(parameters) + self._publish_channel = self._publish_connection.channel() + # 声明队列(持久化) + self._publish_channel.queue_declare(queue=TASK_QUEUE_NAME, durable=True) + logger.info("RabbitMQ 发布连接成功") + except Exception as e: + logger.error(f"RabbitMQ 连接失败: {e}") + raise + + def connect_consume(self): + """创建用于消费的独立连接(线程安全)""" + try: + parameters = self._get_connection_parameters() + connection = pika.BlockingConnection(parameters) + channel = connection.channel() + # 声明队列(持久化) + channel.queue_declare(queue=TASK_QUEUE_NAME, durable=True) + return connection, channel + except Exception as e: + logger.error(f"创建消费连接失败: {e}") + raise + + def _ensure_publish_connection(self): + """确保发布连接可用""" + if not self._publish_connection or self._publish_connection.is_closed: + parameters = self._get_connection_parameters() + self._publish_connection = pika.BlockingConnection(parameters) + self._publish_channel = self._publish_connection.channel() + self._publish_channel.queue_declare(queue=TASK_QUEUE_NAME, durable=True) + + def publish_task(self, task_type: str, task_data: dict): + """发布任务到队列(线程安全)""" + with self._publish_lock: + try: + self._ensure_publish_connection() + + message = { + "task_type": task_type, + "task_data": task_data, + } + self._publish_channel.basic_publish( + exchange="", + routing_key=TASK_QUEUE_NAME, + body=json.dumps(message), + properties=pika.BasicProperties( + delivery_mode=2, # 使消息持久化 + ), + ) + logger.info( + f"任务已入队: task_type={task_type}, task_id={task_data.get('task_id')}" + ) + except Exception as e: + logger.error(f"发布任务到队列失败: {e}") + # 尝试重新连接 + try: + self._ensure_publish_connection() + # 重试一次 + self._publish_channel.basic_publish( + exchange="", + routing_key=TASK_QUEUE_NAME, + body=json.dumps(message), + properties=pika.BasicProperties(delivery_mode=2), + ) + logger.info( + f"任务已入队(重试成功): task_type={task_type}, task_id={task_data.get('task_id')}" + ) + except Exception as retry_err: + logger.error(f"发布任务重试失败: {retry_err}") + raise + + def close(self): + """关闭所有连接""" + try: + if self._publish_connection and not self._publish_connection.is_closed: + self._publish_connection.close() + except Exception: + pass + + +# 全局 RabbitMQ 管理器实例 +rabbitmq_manager = RabbitMQManager() + +# 任务执行线程池(支持5个任务并发执行) +task_executor = futures.ThreadPoolExecutor(max_workers=5) + + +def execute_task(task_type: str, task_data: dict): + """执行任务的通用函数""" + task_id = task_data.get("task_id") + if not task_id: + logger.error("任务 ID 不能为空") + return + + try: + logger.info(f"开始执行任务 {task_id}: {task_type}") + mark_task_started(task_id) + log_task_event(task_id, "INFO", f"开始执行任务: {task_type}") + + if task_type == "collect_binance": + csv_path = os.path.join( + os.path.dirname(__file__), "ETHUSDT-trades-2025-09.csv" + ) + collect_binance.collect_binance( + task_id=task_id, + csv_path=csv_path, + import_percentage=task_data.get("import_percentage", 100), + chunk_size=task_data.get("chunk_size", 1000000), + ) + success_msg = "Binance 数据导入完成" + log_msg = "收集 Binance 数据完成" + + elif task_type == "collect_binance_by_date": + collect_binance.collect_binance_by_date( + task_id=task_id, + start_ts=task_data.get("start_ts"), + end_ts=task_data.get("end_ts"), + ) + success_msg = "Binance 数据按日期收集完成" + log_msg = "按日期收集 Binance 数据完成" + + elif task_type == "collect_uniswap": + collect_uniswap.collect_uniswap( + task_id=task_id, + pool_address=task_data.get("pool_address"), + start_ts=task_data.get("start_ts"), + end_ts=task_data.get("end_ts"), + ) + success_msg = "Uniswap 数据采集完成" + log_msg = "收集 Uniswap 数据完成" + + elif task_type == "process_prices": + # 将时间戳转换为日期字符串 + start_date_str = None + end_date_str = None + if task_data.get("start_date"): + dt = datetime.datetime.fromtimestamp( + task_data.get("start_date"), tz=datetime.timezone.utc + ) + start_date_str = dt.isoformat() + if task_data.get("end_date"): + dt = datetime.datetime.fromtimestamp( + task_data.get("end_date"), tz=datetime.timezone.utc + ) + end_date_str = dt.isoformat() + + kwargs = { + "aggregation_interval": task_data.get("aggregation_interval", "minute"), + "overwrite": task_data.get("overwrite", False), + "start_date": start_date_str, + "end_date": end_date_str, + } + # 合并 db_overrides + db_overrides = task_data.get("db_overrides", {}) + if db_overrides: + kwargs.update(db_overrides) + + process_prices.run_process_prices(task_id=task_id, **kwargs) + success_msg = "价格数据处理完成" + log_msg = "处理价格数据完成" + + elif task_type == "analyse": + # 解析策略参数 + strategy_params = task_data.get("strategy_params", {}) + kwargs = { + "strategy": strategy_params, + "batch_id": task_data.get("batch_id"), + "overwrite": task_data.get("overwrite", False), + } + analyse.run_analyse(task_id=task_id, config_json=json.dumps(kwargs)) + success_msg = "数据分析完成" + log_msg = "分析数据完成" + + else: + raise ValueError(f"未知的任务类型: {task_type}") + + # 检查任务是否被取消 + if check_task(task_id): + logger.info(f"任务 {task_id} 已被取消") + log_task_event(task_id, "INFO", "任务被取消") + mark_task_finished(task_id, TaskStatus.CANCELLED, "任务被取消") + else: + logger.info(f"任务 {task_id} 执行成功: {task_type}") + log_task_event(task_id, "INFO", log_msg) + mark_task_finished(task_id, TaskStatus.SUCCESS, success_msg) + + except Exception as e: + logger.error(f"任务 {task_id} 执行失败: {e}") + log_task_event(task_id, "ERROR", f"任务执行失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, str(e)) + + +def consume_tasks(): + """从队列中消费任务并并发执行(支持5个任务并发)""" + + def callback(ch, method, properties, body): + try: + message = json.loads(body) + task_type = message.get("task_type") + task_data = message.get("task_data") + + if not task_type or not task_data: + logger.error(f"无效的任务消息: {message}") + ch.basic_ack(delivery_tag=method.delivery_tag) + return + + # 在线程池中异步执行任务 + # 注意:我们在任务提交到线程池后立即确认消息 + # 因为任务状态已经在数据库中记录,即使任务失败也不应该重复处理 + task_executor.submit(execute_task, task_type, task_data) + + # 立即确认消息(任务已提交到线程池,状态会在数据库中记录) + ch.basic_ack(delivery_tag=method.delivery_tag) + + except Exception as e: + logger.error(f"处理队列消息失败: {e}") + # 发生异常时立即确认消息,避免重复处理 + try: + ch.basic_ack(delivery_tag=method.delivery_tag) + except Exception: + pass + + # 持续监听队列 + while True: + connection = None + channel = None + try: + # 创建独立的消费连接(与发布连接分离,保证线程安全) + connection, channel = rabbitmq_manager.connect_consume() + + # 设置 QoS,每个消费者可以预取5个任务,支持并发执行 + channel.basic_qos(prefetch_count=5) + # 开始消费 + channel.basic_consume( + queue=TASK_QUEUE_NAME, + on_message_callback=callback, + ) + + logger.info("开始监听任务队列(支持5个任务并发执行)...") + # start_consuming() 会阻塞,直到连接关闭或出现异常 + channel.start_consuming() + except (pika.exceptions.ConnectionClosed, pika.exceptions.ChannelClosed) as e: + logger.warning(f"RabbitMQ 消费连接关闭: {e},5秒后重试...") + try: + if connection and not connection.is_closed: + connection.close() + except Exception: + pass + time.sleep(5) + except Exception as e: + logger.error(f"消费任务时出错: {e},5秒后重试...") + try: + if connection and not connection.is_closed: + connection.close() + except Exception: + pass + time.sleep(5) + + class TaskService(TaskServiceServicer): """实现 TaskService 的 gRPC 服务""" def CollectBinance(self, request, context): """ 收集币安数据 - 在后台线程中执行任务,立即返回运行状态 + 将任务放入队列,立即返回等待状态 """ task_id = request.task_id import_percentage = request.import_percentage @@ -117,45 +431,36 @@ def CollectBinance(self, request, context): f"import_percentage={import_percentage}, chunk_size={chunk_size}" ) - # 获取 CSV 文件路径(相对于 data 目录) - csv_path = os.path.join(os.path.dirname(__file__), "ETHUSDT-trades-2025-09.csv") - - def run_task(): - """在后台线程中执行任务""" - try: - logger.info(f"开始执行任务 {task_id}: 收集币安数据") - mark_task_started(task_id) - log_task_event(task_id, "INFO", "开始收集 Binance 数据") - collect_binance.collect_binance( - task_id=task_id, - csv_path=csv_path, - import_percentage=import_percentage, - chunk_size=chunk_size, - ) - logger.info(f"任务 {task_id} 执行成功: 收集币安数据") - log_task_event(task_id, "INFO", "收集 Binance 数据完成") - mark_task_finished( - task_id, TaskStatus.TASK_STATUS_SUCCESS, "Binance 数据导入完成" - ) - except Exception as e: - logger.error(f"任务 {task_id} 执行失败: {e}") - log_task_event(task_id, "ERROR", f"收集 Binance 数据失败: {e}") - mark_task_finished(task_id, TaskStatus.TASK_STATUS_FAILED, str(e)) - - # 在后台线程中启动任务 - thread = threading.Thread(target=run_task, daemon=True) - thread.start() + # 将任务状态设置为 WAIT + mark_task_waiting(task_id) + log_task_event(task_id, "INFO", "任务已入队,等待执行") + + # 将任务放入队列 + task_data = { + "task_id": task_id, + "import_percentage": import_percentage, + "chunk_size": chunk_size, + } + try: + rabbitmq_manager.publish_task("collect_binance", task_data) + except Exception as e: + logger.error(f"将任务放入队列失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, f"任务入队失败: {e}") + return TaskResponse( + task_id=task_id, + status=TaskStatus.FAILED, + ) - # 立即返回运行状态 + # 返回等待状态 return TaskResponse( task_id=task_id, - status=TaskStatus.TASK_STATUS_RUNNING, + status=TaskStatus.WAIT, ) def CollectBinanceByDate(self, request, context): """ 按日期收集币安数据 - 在后台线程中执行任务,立即返回运行状态 + 将任务放入队列,立即返回等待状态 """ task_id = request.task_id start_ts = request.start_ts @@ -166,88 +471,36 @@ def CollectBinanceByDate(self, request, context): f"start_ts={start_ts}, end_ts={end_ts}" ) - def run_task(): - """在后台线程中执行任务""" - try: - logger.info(f"开始执行任务 {task_id}: 按日期收集币安数据") - mark_task_started(task_id) - log_task_event(task_id, "INFO", "开始按日期收集 Binance 数据") - collect_binance.collect_binance_by_date( - task_id=task_id, - start_ts=start_ts, - end_ts=end_ts, - ) - logger.info(f"任务 {task_id} 执行成功: 按日期收集币安数据") - log_task_event(task_id, "INFO", "按日期收集 Binance 数据完成") - mark_task_finished( - task_id, - TaskStatus.TASK_STATUS_SUCCESS, - "Binance 数据按日期收集完成", - ) - except Exception as e: - logger.error(f"任务 {task_id} 执行失败: {e}") - log_task_event(task_id, "ERROR", f"按日期收集 Binance 数据失败: {e}") - mark_task_finished(task_id, TaskStatus.TASK_STATUS_FAILED, str(e)) - - # 在后台线程中启动任务 - thread = threading.Thread(target=run_task, daemon=True) - thread.start() + # 将任务状态设置为 WAIT + mark_task_waiting(task_id) + log_task_event(task_id, "INFO", "任务已入队,等待执行") + + # 将任务放入队列 + task_data = { + "task_id": task_id, + "start_ts": start_ts, + "end_ts": end_ts, + } + try: + rabbitmq_manager.publish_task("collect_binance_by_date", task_data) + except Exception as e: + logger.error(f"将任务放入队列失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, f"任务入队失败: {e}") + return TaskResponse( + task_id=task_id, + status=TaskStatus.FAILED, + ) - # 立即返回运行状态 + # 返回等待状态 return TaskResponse( task_id=task_id, - status=TaskStatus.TASK_STATUS_RUNNING, + status=TaskStatus.WAIT, ) - def Analyse(self, request, context): - """ - 执行套利分析任务 - """ - task_id = request.task_id - batch_id = request.batch_id - overwrite = request.overwrite - strategy_json = request.strategy_json - - logger.info( - "收到套利分析请求: task_id=%s batch_id=%s overwrite=%s", - task_id, - batch_id, - overwrite, - ) - - def run_task(): - try: - mark_task_started(task_id) - log_task_event(task_id, "INFO", "套利分析开始执行") - config = { - "batch_id": batch_id, - "overwrite": overwrite, - } - if strategy_json: - try: - config["strategy"] = json.loads(strategy_json) - except json.JSONDecodeError as exc: - logger.warning("解析 strategy_json 失败: %s", exc) - analyse.run_analyse(task_id=task_id, config_json=json.dumps(config)) - logger.info("套利分析任务 %s 执行成功", task_id) - log_task_event(task_id, "INFO", "套利分析完成") - mark_task_finished( - task_id, TaskStatus.TASK_STATUS_SUCCESS, "套利分析完成" - ) - except Exception as exc: - logger.error("套利分析任务 %s 执行失败: %s", task_id, exc) - log_task_event(task_id, "ERROR", f"套利分析失败: {exc}") - mark_task_finished(task_id, TaskStatus.TASK_STATUS_FAILED, str(exc)) - - thread = threading.Thread(target=run_task, daemon=True) - thread.start() - - return TaskResponse(task_id=task_id, status=TaskStatus.TASK_STATUS_RUNNING) - def CollectUniswap(self, request, context): """ 收集 Uniswap 数据 - 在后台线程中执行任务,立即返回运行状态 + 将任务放入队列,立即返回等待状态 """ task_id = request.task_id pool_address = request.pool_address @@ -259,42 +512,37 @@ def CollectUniswap(self, request, context): f"pool_address={pool_address}, start_ts={start_ts}, end_ts={end_ts}" ) - def run_task(): - """在后台线程中执行任务""" - try: - logger.info(f"开始执行任务 {task_id}: 收集 Uniswap 数据") - mark_task_started(task_id) - log_task_event(task_id, "INFO", "开始收集 Uniswap 数据") - collect_uniswap.collect_uniswap( - task_id=task_id, - pool_address=pool_address, - start_ts=start_ts, - end_ts=end_ts, - ) - logger.info(f"任务 {task_id} 执行成功: 收集 Uniswap 数据") - log_task_event(task_id, "INFO", "收集 Uniswap 数据完成") - mark_task_finished( - task_id, TaskStatus.TASK_STATUS_SUCCESS, "Uniswap 数据采集完成" - ) - except Exception as e: - logger.error(f"任务 {task_id} 执行失败: {e}") - log_task_event(task_id, "ERROR", f"收集 Uniswap 数据失败: {e}") - mark_task_finished(task_id, TaskStatus.TASK_STATUS_FAILED, str(e)) - - # 在后台线程中启动任务 - thread = threading.Thread(target=run_task, daemon=True) - thread.start() + # 将任务状态设置为 WAIT + mark_task_waiting(task_id) + log_task_event(task_id, "INFO", "任务已入队,等待执行") + + # 将任务放入队列 + task_data = { + "task_id": task_id, + "pool_address": pool_address, + "start_ts": start_ts, + "end_ts": end_ts, + } + try: + rabbitmq_manager.publish_task("collect_uniswap", task_data) + except Exception as e: + logger.error(f"将任务放入队列失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, f"任务入队失败: {e}") + return TaskResponse( + task_id=task_id, + status=TaskStatus.FAILED, + ) - # 立即返回运行状态 + # 返回等待状态 return TaskResponse( task_id=task_id, - status=TaskStatus.TASK_STATUS_RUNNING, + status=TaskStatus.WAIT, ) def ProcessPrices(self, request, context): """ 处理价格数据 - 在后台线程中执行任务,立即返回运行状态 + 将任务放入队列,立即返回等待状态 """ task_id = request.task_id start_date = request.start_date @@ -309,61 +557,39 @@ def ProcessPrices(self, request, context): f"aggregation_interval={aggregation_interval}, overwrite={overwrite}" ) - def run_task(): - """在后台线程中执行任务""" - try: - logger.info(f"开始执行任务 {task_id}: 处理价格数据") - - # 将 int32 时间戳转换为日期字符串 - # 假设 start_date 和 end_date 是 Unix 时间戳(秒) - start_date_str = None - end_date_str = None - - if start_date: - dt = datetime.datetime.fromtimestamp( - start_date, tz=datetime.timezone.utc - ) - start_date_str = dt.isoformat() - - if end_date: - dt = datetime.datetime.fromtimestamp( - end_date, tz=datetime.timezone.utc - ) - end_date_str = dt.isoformat() - - # 准备参数 - kwargs = { - "aggregation_interval": ( - aggregation_interval if aggregation_interval else "minute" - ), - "overwrite": overwrite, - "start_date": start_date_str, - "end_date": end_date_str, - } - - # 合并 db_overrides(如果有) - if db_overrides: - kwargs.update(db_overrides) - - process_prices.run_process_prices(task_id=task_id, **kwargs) - logger.info(f"任务 {task_id} 执行成功: 处理价格数据") - except Exception as e: - logger.error(f"任务 {task_id} 执行失败: {e}") - - # 在后台线程中启动任务 - thread = threading.Thread(target=run_task, daemon=True) - thread.start() + # 将任务状态设置为 WAIT + mark_task_waiting(task_id) + log_task_event(task_id, "INFO", "任务已入队,等待执行") + + # 将任务放入队列 + task_data = { + "task_id": task_id, + "start_date": start_date, + "end_date": end_date, + "aggregation_interval": aggregation_interval, + "overwrite": overwrite, + "db_overrides": db_overrides, + } + try: + rabbitmq_manager.publish_task("process_prices", task_data) + except Exception as e: + logger.error(f"将任务放入队列失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, f"任务入队失败: {e}") + return TaskResponse( + task_id=task_id, + status=TaskStatus.FAILED, + ) - # 立即返回运行状态 + # 返回等待状态 return TaskResponse( task_id=task_id, - status=TaskStatus.TASK_STATUS_RUNNING, + status=TaskStatus.WAIT, ) def Analyse(self, request, context): """ 分析数据 - 在后台线程中执行任务,立即返回运行状态 + 将任务放入队列,立即返回等待状态 """ task_id = request.task_id batch_id = request.batch_id @@ -375,50 +601,62 @@ def Analyse(self, request, context): f"batch_id={batch_id}, overwrite={overwrite}" ) - def run_task(): - """在后台线程中执行任务""" + # 解析策略 JSON + strategy_params = {} + if strategy_json: try: - logger.info(f"开始执行任务 {task_id}: 分析数据") - - # 解析策略 JSON - strategy_params = {} - if strategy_json: - try: - strategy_params = json.loads(strategy_json) - except json.JSONDecodeError as e: - logger.error(f"解析策略 JSON 失败: {e}") - raise ValueError(f"无效的策略 JSON: {e}") - - # 准备参数:先合并所有策略参数,然后添加控制参数 - kwargs = {} - # 合并策略参数(所有策略参数都可以传入) - kwargs["strategy"] = strategy_params - # 添加控制参数 - kwargs["batch_id"] = batch_id - kwargs["overwrite"] = ( - overwrite # overwrite=True 时重建表,overwrite=False 时追加数据 + strategy_params = json.loads(strategy_json) + except json.JSONDecodeError as e: + logger.error(f"解析策略 JSON 失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, f"无效的策略 JSON: {e}") + return TaskResponse( + task_id=task_id, + status=TaskStatus.FAILED, ) - analyse.run_analyse(task_id=task_id, config_json=json.dumps(kwargs)) - - analyse.run_analyse(task_id=task_id, config_json=json.dumps(kwargs)) - logger.info(f"任务 {task_id} 执行成功: 分析数据") - except Exception as e: - logger.error(f"任务 {task_id} 执行失败: {e}") - - # 在后台线程中启动任务 - thread = threading.Thread(target=run_task, daemon=True) - thread.start() + # 将任务状态设置为 WAIT + mark_task_waiting(task_id) + log_task_event(task_id, "INFO", "任务已入队,等待执行") + + # 将任务放入队列 + task_data = { + "task_id": task_id, + "batch_id": batch_id, + "overwrite": overwrite, + "strategy_params": strategy_params, + } + try: + rabbitmq_manager.publish_task("analyse", task_data) + except Exception as e: + logger.error(f"将任务放入队列失败: {e}") + mark_task_finished(task_id, TaskStatus.FAILED, f"任务入队失败: {e}") + return TaskResponse( + task_id=task_id, + status=TaskStatus.FAILED, + ) - # 立即返回运行状态 + # 返回等待状态 return TaskResponse( task_id=task_id, - status=TaskStatus.TASK_STATUS_RUNNING, + status=TaskStatus.WAIT, ) def serve(): - """启动 gRPC 服务器""" + """启动 gRPC 服务器和 RabbitMQ 消费者""" + # 初始化 RabbitMQ 连接 + try: + rabbitmq_manager.connect() + logger.info("RabbitMQ 连接初始化成功") + except Exception as e: + logger.error(f"RabbitMQ 连接初始化失败: {e}") + logger.warning("继续启动 gRPC 服务器,但任务队列功能可能不可用") + + # 启动消费者线程(后台线程) + consumer_thread = threading.Thread(target=consume_tasks, daemon=True) + consumer_thread.start() + logger.info("任务消费者线程已启动") + # 创建线程池执行器 server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) @@ -436,7 +674,11 @@ def serve(): try: server.wait_for_termination() except KeyboardInterrupt: + logger.info("正在关闭服务器...") server.stop(0) + # 关闭任务执行线程池 + task_executor.shutdown(wait=True) + rabbitmq_manager.close() print("gRPC Server stopped.") diff --git a/data/task_stress.json b/data/task_stress.json new file mode 100644 index 0000000..b3c0d8c --- /dev/null +++ b/data/task_stress.json @@ -0,0 +1,15 @@ +{ + "proto": "../protos/task.proto", + "call": "task.v1.TaskService.ProcessPrices", + "total": 200, + "concurrency": 100, + "insecure": true, + "data": { + "task_id": "5d6f9837-c4ac-4a88-b67e-cd1cb1e251fd", + "start_date": 1756684800, + "end_date": 1757030400, + "aggregation_interval": "1m", + "overwrite": true, + "db_overrides": {} + } +} \ No newline at end of file diff --git a/data/test/test_analyse.py b/data/test/test_analyse.py index 95b2077..e72c44d 100644 --- a/data/test/test_analyse.py +++ b/data/test/test_analyse.py @@ -627,3 +627,321 @@ def test_save_results_rollback_on_error(self): # 验证回滚被调用(如果发生异常) # 注意:由于使用了 context manager,可能不会调用 rollback + + +class TestParseTimestamp: + """ + 测试时间戳解析函数 + """ + + def test_parse_timestamp_valid(self): + """ + 测试:解析有效的时间戳字符串 + """ + from block_chain.analyse import _parse_timestamp + + result = _parse_timestamp("2025-01-01 10:00:00") + assert isinstance(result, pd.Timestamp) + assert result.tz is not None + + def test_parse_timestamp_empty(self): + """ + 测试:空字符串 + """ + from block_chain.analyse import _parse_timestamp + + result = _parse_timestamp("") + assert result is None + + def test_parse_timestamp_none(self): + """ + 测试:None值 + """ + from block_chain.analyse import _parse_timestamp + + result = _parse_timestamp(None) + assert result is None + + +class TestEnsureBatchExists: + """ + 测试确保批次存在函数 + """ + + def test_ensure_batch_exists_batch_exists(self): + """ + 测试:批次已存在 + """ + from block_chain.analyse import ensure_batch_exists + + mock_conn = MagicMock() + mock_cur = MagicMock() + mock_conn.cursor.return_value.__enter__ = lambda x: mock_cur + mock_conn.cursor.return_value.__exit__ = lambda *args: None + mock_cur.fetchone.return_value = (1,) # 批次存在 + + ensure_batch_exists(mock_conn, 1) + + # 应该查询批次,但不创建 + assert mock_cur.execute.called + # 如果批次已存在,不会调用commit(因为只在创建批次时才commit) + # 但为了代码一致性,可能会调用commit,所以不强制检查 + + def test_ensure_batch_exists_batch_not_exists(self): + """ + 测试:批次不存在,自动创建 + """ + from block_chain.analyse import ensure_batch_exists + + mock_conn = MagicMock() + mock_cur = MagicMock() + mock_conn.cursor.return_value.__enter__ = lambda x: mock_cur + mock_conn.cursor.return_value.__exit__ = lambda *args: None + mock_cur.fetchone.return_value = None # 批次不存在 + + ensure_batch_exists(mock_conn, 1) + + # 应该创建批次 + assert mock_cur.execute.call_count >= 2 # SELECT + INSERT + assert mock_conn.commit.called + + def test_ensure_batch_exists_zero_batch_id(self): + """ + 测试:batch_id为0或None,不执行任何操作 + """ + from block_chain.analyse import ensure_batch_exists + + mock_conn = MagicMock() + mock_cur = MagicMock() + mock_conn.cursor.return_value.__enter__ = lambda x: mock_cur + mock_conn.cursor.return_value.__exit__ = lambda *args: None + + ensure_batch_exists(mock_conn, 0) + ensure_batch_exists(mock_conn, None) + + # 不应该执行任何操作 + assert not mock_cur.execute.called + + +class TestRunAnalyse: + """ + 测试主分析函数 + """ + + @patch("block_chain.analyse.fetch_price_pairs") + @patch("block_chain.analyse.analyze_opportunities") + @patch("block_chain.analyse.save_results") + @patch("block_chain.analyse.ensure_batch_exists") + @patch("block_chain.analyse.check_task", return_value=False) + @patch("block_chain.analyse.update_task_status") + @patch("block_chain.analyse.psycopg2.connect") + def test_run_analyse_success( + self, + mock_connect, + mock_update_status, + mock_check_task, + mock_ensure_batch, + mock_save_results, + mock_analyze, + mock_fetch, + ): + """ + 测试:成功运行分析 + """ + from block_chain.analyse import run_analyse + + mock_conn = MagicMock() + mock_connect.return_value = mock_conn + mock_fetch.return_value = [] + mock_analyze.return_value = [] + + config_json = '{"strategy": {}, "batch_id": 1, "overwrite": false}' + + run_analyse("test_task", config_json) + + mock_ensure_batch.assert_called_once() + mock_fetch.assert_called_once() + mock_analyze.assert_called_once() + mock_save_results.assert_called_once() + mock_update_status.assert_called_once_with("test_task", 1) + + @patch("block_chain.analyse.update_task_status") + @patch("block_chain.analyse.psycopg2.connect") + def test_run_analyse_invalid_time_range( + self, + mock_connect, + mock_update_status, + ): + """ + 测试:无效的时间范围(start > end) + """ + from block_chain.analyse import run_analyse + + mock_conn = MagicMock() + mock_connect.return_value = mock_conn + + config_json = ( + '{"strategy": {"start": "2025-01-02", "end": "2025-01-01"}, ' + '"batch_id": 1, "overwrite": false}' + ) + + run_analyse("test_task", config_json) + + mock_update_status.assert_called_once_with("test_task", 2) + + @patch("block_chain.analyse.fetch_price_pairs") + @patch("block_chain.analyse.ensure_batch_exists") + @patch("block_chain.analyse.check_task", return_value=False) + @patch("block_chain.analyse.update_task_status") + @patch("block_chain.analyse.psycopg2.connect") + def test_run_analyse_exception_handling( + self, + mock_connect, + mock_update_status, + mock_check_task, + mock_ensure_batch, + mock_fetch, + ): + """ + 测试:异常处理和回滚 + """ + from block_chain.analyse import run_analyse + + mock_conn = MagicMock() + mock_connect.return_value = mock_conn + mock_fetch.side_effect = Exception("Database error") + + config_json = '{"strategy": {}, "batch_id": 1, "overwrite": false}' + + with pytest.raises(Exception, match="Database error"): + run_analyse("test_task", config_json) + + mock_conn.rollback.assert_called_once() + mock_update_status.assert_called_once_with("test_task", 2) + + @patch("block_chain.analyse.fetch_price_pairs") + @patch("block_chain.analyse.analyze_opportunities") + @patch("block_chain.analyse.save_results") + @patch("block_chain.analyse.ensure_batch_exists") + @patch("block_chain.analyse.check_task") + @patch("block_chain.analyse.update_task_status") + @patch("block_chain.analyse.psycopg2.connect") + def test_run_analyse_task_cancelled( + self, + mock_connect, + mock_update_status, + mock_check_task, + mock_ensure_batch, + mock_save_results, + mock_analyze, + mock_fetch, + ): + """ + 测试:任务被取消 + """ + from block_chain.analyse import run_analyse + + mock_conn = MagicMock() + mock_connect.return_value = mock_conn + mock_fetch.return_value = [] + mock_analyze.return_value = [] + # 根据代码逻辑,check_task在最后才被调用一次(在else块中) + # 如果返回True,不会调用update_task_status + # 所以我们需要让check_task在第一次(也是唯一一次)调用时返回True + mock_check_task.return_value = True + + config_json = '{"strategy": {}, "batch_id": 1, "overwrite": false}' + + run_analyse("test_task", config_json) + + # 由于check_task返回True,update_task_status不应该被调用 + mock_update_status.assert_not_called() + + @patch("block_chain.analyse.fetch_price_pairs") + @patch("block_chain.analyse.analyze_opportunities") + @patch("block_chain.analyse.save_results") + @patch("block_chain.analyse.ensure_batch_exists") + @patch("block_chain.analyse.check_task", return_value=False) + @patch("block_chain.analyse.update_task_status") + @patch("block_chain.analyse.psycopg2.connect") + def test_run_analyse_with_custom_strategy( + self, + mock_connect, + mock_update_status, + mock_check_task, + mock_ensure_batch, + mock_save_results, + mock_analyze, + mock_fetch, + ): + """ + 测试:使用自定义策略配置 + """ + from block_chain.analyse import run_analyse + + mock_conn = MagicMock() + mock_connect.return_value = mock_conn + mock_fetch.return_value = [] + mock_analyze.return_value = [] + + config_json = ( + '{"strategy": {"initial_investment": 50000, "profit_threshold": 5}, ' + '"batch_id": 2, "overwrite": true, "experiment_id": 123}' + ) + + run_analyse("test_task", config_json) + + # 验证analyze_opportunities被调用,并且策略参数被传递 + assert mock_analyze.called + call_args = mock_analyze.call_args + strategy = call_args[0][1] + assert strategy["initial_investment"] == 50000 + assert strategy["profit_threshold"] == 5 + + @patch("block_chain.analyse.fetch_price_pairs") + @patch("block_chain.analyse.analyze_opportunities") + @patch("block_chain.analyse.save_results") + @patch("block_chain.analyse.ensure_batch_exists") + @patch("block_chain.analyse.check_task", return_value=False) + @patch("block_chain.analyse.update_task_status") + @patch("block_chain.analyse.psycopg2.connect") + def test_run_analyse_with_time_range( + self, + mock_connect, + mock_update_status, + mock_check_task, + mock_ensure_batch, + mock_save_results, + mock_analyze, + mock_fetch, + ): + """ + 测试:使用时间范围 + """ + from block_chain.analyse import run_analyse + + mock_conn = MagicMock() + mock_connect.return_value = mock_conn + mock_fetch.return_value = [] + mock_analyze.return_value = [] + + config_json = ( + '{"strategy": {"start": "2025-01-01", "end": "2025-01-02"}, ' + '"batch_id": 1, "overwrite": false}' + ) + + run_analyse("test_task", config_json) + + # 验证fetch_price_pairs被调用,并且时间参数被传递 + assert mock_fetch.called + call_args = mock_fetch.call_args + # 参数可能是位置参数或关键字参数 + if len(call_args) > 1 and "start_time" in call_args[1]: + start_time = call_args[1]["start_time"] + end_time = call_args[1]["end_time"] + else: + # 可能是位置参数 + start_time = call_args[0][2] if len(call_args[0]) > 2 else None + end_time = call_args[0][3] if len(call_args[0]) > 3 else None + assert start_time is not None + assert end_time is not None diff --git a/data/test/test_collect_binance.py b/data/test/test_collect_binance.py index f12ac5c..dc07bed 100644 --- a/data/test/test_collect_binance.py +++ b/data/test/test_collect_binance.py @@ -232,9 +232,7 @@ def test_process_chunk_handles_db_error(self, sample_chunk, mock_db_connection): None, ) # 验证任务状态被更新为失败 - mock_update_status.assert_called_once_with( - "test_task", "FAILED" - ) + mock_update_status.assert_called_once_with("test_task", "FAILED") def test_process_chunk_stops_at_target_rows(self, mock_db_connection): """ @@ -400,7 +398,7 @@ def test_import_data_to_database_success( mock_read_csv.return_value = [sample_chunk] # process_chunk 会修改 rows_counter,所以我们需要让它实际执行 - def side_effect(task_id, chunk, idx, counter, target): + def side_effect(task_id, chunk, idx, counter, target, conn=None): counter[0] += len(chunk) counter[1] += len(chunk) return (True, len(chunk), len(chunk), False) @@ -446,7 +444,7 @@ def test_import_data_to_database_with_target_rows( mock_read_csv.return_value = [sample_chunk] # 第一个chunk达到目标行数 - def side_effect(task_id, chunk, idx, counter, target): + def side_effect(task_id, chunk, idx, counter, target, conn=None): counter[0] += len(chunk) counter[1] += len(chunk) should_stop = counter[0] >= target if target else False @@ -537,7 +535,7 @@ def test_import_data_to_database_multiple_chunks( ) mock_read_csv.return_value = [chunk1, chunk2] - def side_effect(task_id, chunk, idx, counter, target): + def side_effect(task_id, chunk, idx, counter, target, conn=None): counter[0] += len(chunk) counter[1] += len(chunk) return (True, len(chunk), len(chunk), False) @@ -596,7 +594,7 @@ def test_import_data_to_database_stops_at_target( mock_read_csv.return_value = [chunk1, chunk2] # 第一个chunk达到目标行数,返回 should_stop=True - def side_effect(task_id, chunk, idx, counter, target): + def side_effect(task_id, chunk, idx, counter, target, conn=None): counter[0] += len(chunk) counter[1] += len(chunk) should_stop = counter[0] >= target if target else False @@ -611,3 +609,566 @@ def side_effect(task_id, chunk, idx, counter, target): assert rows_counter[0] >= 3 # 第二个chunk不应该被处理 assert mock_process_chunk.call_count == 1 + + def test_process_chunk_with_conn_parameter( + self, sample_binance_chunk, mock_db_connection + ): + """ + 测试:使用提供的数据库连接参数 + """ + mock_conn, mock_cursor = mock_db_connection + rows_counter = [0, 0] + + success, rows_processed, rows_imported, should_stop = process_chunk( + "test_task", + sample_binance_chunk, + 0, + rows_counter, + None, + conn=mock_conn, + ) + + assert success is True + assert rows_processed == len(sample_binance_chunk) + assert rows_imported == len(sample_binance_chunk) + # 验证使用了提供的连接,而不是创建新连接 + mock_cursor.copy_expert.assert_called_once() + # 不应该调用 commit(由调用者控制事务) + assert not mock_conn.commit.called + + +class TestCalcTargetRows: + """ + 测试计算目标行数函数 + """ + + def test_calc_target_rows_with_percentage(self): + """ + 测试:使用百分比计算目标行数 + """ + from block_chain.collect_binance import _calc_target_rows + + result = _calc_target_rows(1000, 50) + assert result == 500 + + result = _calc_target_rows(1000, 100) + assert result == 1000 + + result = _calc_target_rows(1000, 101) # 超过100% + assert result == 1000 + + def test_calc_target_rows_with_none(self): + """ + 测试:total_lines为None的情况 + """ + from block_chain.collect_binance import _calc_target_rows + + result = _calc_target_rows(None, 50) + assert result is None + + +class TestCollectBinance: + """ + 测试主收集函数 + """ + + @patch("block_chain.collect_binance.count_lines", return_value=1000) + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + @patch("block_chain.collect_binance.update_task_status") + def test_collect_binance_success( + self, + mock_update_status, + mock_connect, + mock_import_data, + mock_check_task, + mock_count_lines, + mock_db_connection, + ): + """ + 测试:成功收集数据 + """ + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_import_data.return_value = [1000, 1000] + + from block_chain.collect_binance import collect_binance + + result = collect_binance("test_task", "test.csv", 100, 1000) + + assert result == 1000 + mock_conn.commit.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "SUCCESS") + + @patch("block_chain.collect_binance.count_lines", return_value=1000) + @patch("block_chain.collect_binance.check_task") + @patch("block_chain.collect_binance.psycopg2.connect") + def test_collect_binance_task_cancelled_before_import( + self, + mock_connect, + mock_check_task, + mock_count_lines, + mock_db_connection, + ): + """ + 测试:在导入前任务被取消 + """ + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + # 第一次检查返回False,第二次返回True(任务被取消) + mock_check_task.side_effect = [False, True] + + from block_chain.collect_binance import collect_binance + + # 需要mock import_data_to_database以避免实际读取文件 + with patch( + "block_chain.collect_binance.import_data_to_database" + ) as mock_import: + result = collect_binance("test_task", "test.csv", 100, 1000) + + assert result == 0 + mock_conn.rollback.assert_called_once() + + @patch("block_chain.collect_binance.count_lines", return_value=1000) + @patch("block_chain.collect_binance.check_task") + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + def test_collect_binance_task_cancelled_after_import( + self, + mock_connect, + mock_import_data, + mock_check_task, + mock_count_lines, + mock_db_connection, + ): + """ + 测试:在导入后任务被取消 + """ + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_import_data.return_value = [1000, 1000] + # 第一次和第二次检查返回False,第三次返回True(任务被取消) + # 注意:在collect_binance中,check_task在commit之后再次被调用 + mock_check_task.side_effect = [False, False, True, True] + + from block_chain.collect_binance import collect_binance + + result = collect_binance("test_task", "test.csv", 100, 1000) + + # 由于在commit之后才检查,所以会返回导入的行数,但不会标记为成功 + assert result == 1000 + mock_conn.commit.assert_called_once() + # 不会回滚,因为已经在commit之后了 + + @patch("block_chain.collect_binance.count_lines", return_value=1000) + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + @patch("block_chain.collect_binance.update_task_status") + def test_collect_binance_exception_handling( + self, + mock_update_status, + mock_connect, + mock_import_data, + mock_check_task, + mock_count_lines, + mock_db_connection, + ): + """ + 测试:异常处理和回滚 + """ + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_import_data.side_effect = Exception("Database error") + + from block_chain.collect_binance import collect_binance + + with pytest.raises(Exception, match="Database error"): + collect_binance("test_task", "test.csv", 100, 1000) + + mock_conn.rollback.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "FAILED") + + @patch("block_chain.collect_binance.count_lines", return_value=1000) + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + @patch("block_chain.collect_binance.update_task_status") + def test_collect_binance_rollback_failure( + self, + mock_update_status, + mock_connect, + mock_import_data, + mock_check_task, + mock_count_lines, + mock_db_connection, + ): + """ + 测试:回滚失败的情况 + """ + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_import_data.side_effect = Exception("Database error") + mock_conn.rollback.side_effect = Exception("Rollback failed") + + from block_chain.collect_binance import collect_binance + + with pytest.raises(Exception, match="Database error"): + collect_binance("test_task", "test.csv", 100, 1000) + + # 应该尝试回滚,即使回滚失败 + mock_conn.rollback.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "FAILED") + + @patch("block_chain.collect_binance.count_lines", return_value=1000) + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + def test_collect_binance_close_connection( + self, + mock_connect, + mock_import_data, + mock_check_task, + mock_count_lines, + mock_db_connection, + ): + """ + 测试:确保连接被关闭 + """ + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_import_data.return_value = [1000, 1000] + + from block_chain.collect_binance import collect_binance + + collect_binance("test_task", "test.csv", 100, 1000) + + mock_conn.close.assert_called_once() + + +class TestDownloadBinanceFile: + """ + 测试下载币安文件函数 + """ + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.requests.get") + @patch("block_chain.collect_binance.tempfile.mkdtemp") + @patch("block_chain.collect_binance.zipfile.ZipFile") + @patch("block_chain.collect_binance.os.remove") + @patch("block_chain.collect_binance.os.rmdir") + @patch("block_chain.collect_binance.os.path.join") + @patch("builtins.open", create=True) + def test_download_binance_file_success( + self, + mock_open, + mock_join, + mock_rmdir, + mock_remove, + mock_zipfile, + mock_mkdtemp, + mock_get, + mock_check_task, + ): + """ + 测试:成功下载文件 + """ + from block_chain.collect_binance import download_binance_file + + mock_mkdtemp.return_value = "/tmp/test_dir" + mock_join.side_effect = lambda *args: "/".join(args) + mock_response = MagicMock() + mock_response.iter_content.return_value = [b"chunk1", b"chunk2"] + mock_response.raise_for_status = Mock() + mock_get.return_value = mock_response + + mock_file = MagicMock() + mock_open.return_value.__enter__ = lambda x: mock_file + mock_open.return_value.__exit__ = lambda *args: None + + mock_zip = MagicMock() + mock_zip.namelist.return_value = ["ETHUSDT-trades-2025-01-01.csv"] + mock_zip.extract = Mock() + mock_zipfile.return_value.__enter__ = lambda x: mock_zip + mock_zipfile.return_value.__exit__ = lambda *args: None + + result = download_binance_file("test_task", "2025-01-01", "ETHUSDT") + + assert result is not None + assert "ETHUSDT-trades-2025-01-01.csv" in result + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.requests.get") + @patch("block_chain.collect_binance.tempfile.mkdtemp") + def test_download_binance_file_request_exception( + self, mock_mkdtemp, mock_get, mock_check_task + ): + """ + 测试:请求异常 + """ + import requests + + from block_chain.collect_binance import download_binance_file + + mock_get.side_effect = requests.exceptions.RequestException("Network error") + + result = download_binance_file("test_task", "2025-01-01", "ETHUSDT") + + assert result is None + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.requests.get") + @patch("block_chain.collect_binance.tempfile.mkdtemp") + @patch("block_chain.collect_binance.zipfile.ZipFile") + @patch("block_chain.collect_binance.os.remove") + @patch("block_chain.collect_binance.os.rmdir") + def test_download_binance_file_no_csv_in_zip( + self, + mock_rmdir, + mock_remove, + mock_zipfile, + mock_mkdtemp, + mock_get, + mock_check_task, + ): + """ + 测试:ZIP文件中没有CSV文件 + """ + from block_chain.collect_binance import download_binance_file + + mock_mkdtemp.return_value = "/tmp/test_dir" + mock_response = MagicMock() + mock_response.iter_content.return_value = [b"chunk1"] + mock_response.raise_for_status = Mock() + mock_get.return_value = mock_response + + mock_zip = MagicMock() + mock_zip.namelist.return_value = [] # 没有CSV文件 + mock_zipfile.return_value.__enter__ = lambda x: mock_zip + mock_zipfile.return_value.__exit__ = lambda *args: None + + result = download_binance_file("test_task", "2025-01-01", "ETHUSDT") + + assert result is None + + @patch("block_chain.collect_binance.check_task") + @patch("block_chain.collect_binance.requests.get") + @patch("block_chain.collect_binance.tempfile.mkdtemp") + @patch("block_chain.collect_binance.os.remove") + @patch("block_chain.collect_binance.os.rmdir") + @patch("block_chain.collect_binance.os.path.join") + @patch("builtins.open", create=True) + def test_download_binance_file_task_cancelled( + self, + mock_open, + mock_join, + mock_rmdir, + mock_remove, + mock_mkdtemp, + mock_get, + mock_check_task, + ): + """ + 测试:下载过程中任务被取消 + """ + from block_chain.collect_binance import download_binance_file + + mock_mkdtemp.return_value = "/tmp/test_dir" + mock_join.side_effect = lambda *args: "/".join(args) + mock_response = MagicMock() + # 第一次检查返回False,第二次返回True(任务被取消) + mock_check_task.side_effect = [False, True] + mock_response.iter_content.return_value = [b"chunk1"] + mock_response.raise_for_status = Mock() + mock_get.return_value = mock_response + + mock_file = MagicMock() + mock_open.return_value.__enter__ = lambda x: mock_file + mock_open.return_value.__exit__ = lambda *args: None + + result = download_binance_file("test_task", "2025-01-01", "ETHUSDT") + + assert result is None + # 验证清理操作被调用(可能在异常处理中) + # 由于任务取消发生在下载过程中,文件可能还未创建,所以remove可能不会被调用 + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.requests.get") + @patch("block_chain.collect_binance.tempfile.mkdtemp") + @patch("block_chain.collect_binance.zipfile.ZipFile") + @patch("block_chain.collect_binance.os.remove") + @patch("block_chain.collect_binance.os.rmdir") + def test_download_binance_file_bad_zip( + self, + mock_rmdir, + mock_remove, + mock_zipfile, + mock_mkdtemp, + mock_get, + mock_check_task, + ): + """ + 测试:ZIP文件损坏 + """ + import zipfile + + from block_chain.collect_binance import download_binance_file + + mock_mkdtemp.return_value = "/tmp/test_dir" + mock_response = MagicMock() + mock_response.iter_content.return_value = [b"chunk1"] + mock_response.raise_for_status = Mock() + mock_get.return_value = mock_response + + mock_zipfile.side_effect = zipfile.BadZipFile("Bad zip file") + + result = download_binance_file("test_task", "2025-01-01", "ETHUSDT") + + assert result is None + + +class TestCollectBinanceByDate: + """ + 测试按日期收集币安数据函数 + """ + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.download_binance_file") + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + @patch("block_chain.collect_binance.update_task_status") + @patch("block_chain.collect_binance.os.remove") + @patch("block_chain.collect_binance.os.path.exists") + @patch("block_chain.collect_binance.os.path.isdir") + @patch("block_chain.collect_binance.os.rmdir") + def test_collect_binance_by_date_success( + self, + mock_rmdir, + mock_isdir, + mock_exists, + mock_remove, + mock_update_status, + mock_connect, + mock_import_data, + mock_download, + mock_check_task, + mock_db_connection, + ): + """ + 测试:成功按日期收集数据 + """ + from block_chain.collect_binance import collect_binance_by_date + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_download.return_value = "/tmp/test.csv" + mock_import_data.return_value = [100, 100] + mock_exists.return_value = True + mock_isdir.return_value = True + + start_ts = int(pd.Timestamp("2025-01-01", tz="UTC").timestamp()) + end_ts = int(pd.Timestamp("2025-01-02", tz="UTC").timestamp()) + + result = collect_binance_by_date("test_task", start_ts, end_ts, "ETHUSDT", 1000) + + # 处理了两天,每天100行,总共200行 + assert result == 200 + mock_conn.commit.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "SUCCESS") + + @patch("block_chain.collect_binance.check_task") + @patch("block_chain.collect_binance.download_binance_file") + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + def test_collect_binance_by_date_task_cancelled( + self, + mock_connect, + mock_import_data, + mock_download, + mock_check_task, + mock_db_connection, + ): + """ + 测试:任务被取消 + """ + from block_chain.collect_binance import collect_binance_by_date + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + # 第一次检查返回False,第二次返回True(任务被取消) + mock_check_task.side_effect = [False, True] + + start_ts = int(pd.Timestamp("2025-01-01", tz="UTC").timestamp()) + end_ts = int(pd.Timestamp("2025-01-02", tz="UTC").timestamp()) + + result = collect_binance_by_date("test_task", start_ts, end_ts, "ETHUSDT", 1000) + + assert result == 0 + mock_conn.rollback.assert_called_once() + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.download_binance_file") + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + @patch("block_chain.collect_binance.update_task_status") + def test_collect_binance_by_date_download_fails( + self, + mock_update_status, + mock_connect, + mock_import_data, + mock_download, + mock_check_task, + mock_db_connection, + ): + """ + 测试:下载失败,跳过该日期 + """ + from block_chain.collect_binance import collect_binance_by_date + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_download.return_value = None # 下载失败 + + start_ts = int(pd.Timestamp("2025-01-01", tz="UTC").timestamp()) + end_ts = int(pd.Timestamp("2025-01-01", tz="UTC").timestamp()) + + result = collect_binance_by_date("test_task", start_ts, end_ts, "ETHUSDT", 1000) + + assert result == 0 + # 应该提交事务(即使没有数据) + mock_conn.commit.assert_called_once() + + @patch("block_chain.collect_binance.check_task", return_value=False) + @patch("block_chain.collect_binance.download_binance_file") + @patch("block_chain.collect_binance.import_data_to_database") + @patch("block_chain.collect_binance.psycopg2.connect") + @patch("block_chain.collect_binance.update_task_status") + def test_collect_binance_by_date_import_fails( + self, + mock_update_status, + mock_connect, + mock_import_data, + mock_download, + mock_check_task, + mock_db_connection, + ): + """ + 测试:导入失败,回滚事务 + """ + from block_chain.collect_binance import collect_binance_by_date + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_download.return_value = "/tmp/test.csv" + mock_import_data.side_effect = Exception("Import error") + + start_ts = int(pd.Timestamp("2025-01-01", tz="UTC").timestamp()) + end_ts = int(pd.Timestamp("2025-01-01", tz="UTC").timestamp()) + + with pytest.raises(Exception, match="Import error"): + collect_binance_by_date("test_task", start_ts, end_ts, "ETHUSDT", 1000) + + # rollback可能被调用多次(一次在异常处理中,一次在finally中) + assert mock_conn.rollback.called + mock_update_status.assert_called_once_with("test_task", "FAILED") diff --git a/data/test/test_collect_uniswap.py b/data/test/test_collect_uniswap.py index 5b17d92..7273906 100644 --- a/data/test/test_collect_uniswap.py +++ b/data/test/test_collect_uniswap.py @@ -287,3 +287,236 @@ def test_process_and_store_uniswap_data_empty_data(self, mock_db_connection): # 应该返回0,不执行数据库操作 assert result == 0 + + @patch("block_chain.collect_uniswap.execute_values") + def test_process_and_store_uniswap_data_with_conn( + self, mock_execute_values, mock_db_connection + ): + """ + 测试:使用提供的数据库连接参数 + """ + mock_conn, mock_cursor = mock_db_connection + + swaps_data = [ + { + "id": "0x1", + "timestamp": "1725187200", + "amount0": "1.0", + "amount1": "3000.0", + "transaction": {"id": "0xtx1", "gasPrice": "50000000000"}, + }, + ] + + result = process_and_store_uniswap_data("test_task", swaps_data, conn=mock_conn) + + # 验证execute_values被调用 + mock_execute_values.assert_called_once() + # 验证返回了记录数量 + assert result == 1 + # 验证使用了提供的连接(通过检查execute_values的调用) + call_args = mock_execute_values.call_args + assert call_args is not None + + +class TestCollectUniswap: + """ + 测试主收集函数 + """ + + @patch("block_chain.collect_uniswap.fetch_all_swaps") + @patch("block_chain.collect_uniswap.check_task", return_value=False) + @patch("block_chain.collect_uniswap.process_and_store_uniswap_data") + @patch("block_chain.collect_uniswap.psycopg2.connect") + @patch("block_chain.collect_uniswap.update_task_status") + def test_collect_uniswap_success( + self, + mock_update_status, + mock_connect, + mock_process_data, + mock_check_task, + mock_fetch_swaps, + mock_db_connection, + ): + """ + 测试:成功收集数据 + """ + from block_chain.collect_uniswap import collect_uniswap + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_fetch_swaps.return_value = [ + { + "id": "0x1", + "timestamp": "1725187200", + "amount0": "1.0", + "amount1": "3000.0", + "transaction": {"id": "0xtx1", "gasPrice": "50000000000"}, + } + ] + mock_process_data.return_value = 1 + + result = collect_uniswap("test_task", "0x123", 1725187200, 1725187260) + + assert result == 1 + mock_conn.commit.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "SUCCESS") + + @patch("block_chain.collect_uniswap.fetch_all_swaps") + @patch("block_chain.collect_uniswap.check_task") + @patch("block_chain.collect_uniswap.psycopg2.connect") + def test_collect_uniswap_task_cancelled_after_fetch( + self, + mock_connect, + mock_check_task, + mock_fetch_swaps, + mock_db_connection, + ): + """ + 测试:获取数据后任务被取消 + """ + from block_chain.collect_uniswap import collect_uniswap + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_fetch_swaps.return_value = [ + { + "id": "0x1", + "timestamp": "1725187200", + "amount0": "1.0", + "amount1": "3000.0", + "transaction": {"id": "0xtx1", "gasPrice": "50000000000"}, + } + ] + # 第一次检查返回False,第二次返回True(任务被取消) + mock_check_task.side_effect = [False, True] + + result = collect_uniswap("test_task", "0x123", 1725187200, 1725187260) + + assert result == 0 + mock_conn.rollback.assert_called_once() + + @patch("block_chain.collect_uniswap.fetch_all_swaps") + @patch("block_chain.collect_uniswap.check_task") + @patch("block_chain.collect_uniswap.process_and_store_uniswap_data") + @patch("block_chain.collect_uniswap.psycopg2.connect") + def test_collect_uniswap_task_cancelled_after_process( + self, + mock_connect, + mock_process_data, + mock_check_task, + mock_fetch_swaps, + mock_db_connection, + ): + """ + 测试:处理数据后任务被取消 + """ + from block_chain.collect_uniswap import collect_uniswap + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_fetch_swaps.return_value = [ + { + "id": "0x1", + "timestamp": "1725187200", + "amount0": "1.0", + "amount1": "3000.0", + "transaction": {"id": "0xtx1", "gasPrice": "50000000000"}, + } + ] + mock_process_data.return_value = 1 + # 第一次和第二次检查返回False,第三次返回True(任务被取消) + # 注意:在collect_uniswap中,check_task在commit之后再次被调用 + mock_check_task.side_effect = [False, False, True, True] + + result = collect_uniswap("test_task", "0x123", 1725187200, 1725187260) + + # 由于在commit之后才检查,所以会返回导入的行数,但不会标记为成功 + assert result == 1 + mock_conn.commit.assert_called_once() + # 不会回滚,因为已经在commit之后了 + + @patch("block_chain.collect_uniswap.fetch_all_swaps") + @patch("block_chain.collect_uniswap.check_task", return_value=False) + @patch("block_chain.collect_uniswap.process_and_store_uniswap_data") + @patch("block_chain.collect_uniswap.psycopg2.connect") + @patch("block_chain.collect_uniswap.update_task_status") + def test_collect_uniswap_exception_handling( + self, + mock_update_status, + mock_connect, + mock_process_data, + mock_check_task, + mock_fetch_swaps, + mock_db_connection, + ): + """ + 测试:异常处理和回滚 + """ + from block_chain.collect_uniswap import collect_uniswap + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_fetch_swaps.side_effect = Exception("Network error") + + result = collect_uniswap("test_task", "0x123", 1725187200, 1725187260) + + assert result == 0 + mock_conn.rollback.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "FAILED") + + @patch("block_chain.collect_uniswap.fetch_all_swaps") + @patch("block_chain.collect_uniswap.check_task", return_value=False) + @patch("block_chain.collect_uniswap.process_and_store_uniswap_data") + @patch("block_chain.collect_uniswap.psycopg2.connect") + @patch("block_chain.collect_uniswap.update_task_status") + def test_collect_uniswap_rollback_failure( + self, + mock_update_status, + mock_connect, + mock_process_data, + mock_check_task, + mock_fetch_swaps, + mock_db_connection, + ): + """ + 测试:回滚失败的情况 + """ + from block_chain.collect_uniswap import collect_uniswap + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_fetch_swaps.side_effect = Exception("Network error") + mock_conn.rollback.side_effect = Exception("Rollback failed") + + result = collect_uniswap("test_task", "0x123", 1725187200, 1725187260) + + assert result == 0 + # 应该尝试回滚,即使回滚失败 + mock_conn.rollback.assert_called_once() + mock_update_status.assert_called_once_with("test_task", "FAILED") + + @patch("block_chain.collect_uniswap.fetch_all_swaps") + @patch("block_chain.collect_uniswap.check_task", return_value=False) + @patch("block_chain.collect_uniswap.process_and_store_uniswap_data") + @patch("block_chain.collect_uniswap.psycopg2.connect") + def test_collect_uniswap_close_connection( + self, + mock_connect, + mock_process_data, + mock_check_task, + mock_fetch_swaps, + mock_db_connection, + ): + """ + 测试:确保连接被关闭 + """ + from block_chain.collect_uniswap import collect_uniswap + + mock_conn, mock_cursor = mock_db_connection + mock_connect.return_value = mock_conn + mock_fetch_swaps.return_value = [] + mock_process_data.return_value = 0 + + collect_uniswap("test_task", "0x123", 1725187200, 1725187260) + + mock_conn.close.assert_called_once() diff --git a/protos/task.proto b/protos/task.proto index 65a6431..82e2c4b 100644 --- a/protos/task.proto +++ b/protos/task.proto @@ -8,7 +8,7 @@ service TaskService { // 收集币安数据 rpc CollectBinance(CollectBinanceRequest) returns (TaskResponse); // 按日期收集币安数据 - rpc CollectBinanceByDate(CollectBinaceByDateRequest) returns (TaskResponse); + rpc CollectBinanceByDate(CollectBinanceByDateRequest) returns (TaskResponse); // 收集Uniswap数据 rpc CollectUniswap(CollectUniswapRequest) returns (TaskResponse); // 处理价格数据 @@ -19,10 +19,11 @@ service TaskService { // TaskStatus 定义任务的可能状态。 enum TaskStatus { - TASK_STATUS_RUNNING = 0; // 任务正在执行中 - TASK_STATUS_SUCCESS = 1; // 任务成功完成 - TASK_STATUS_FAILED = 2; // 任务执行失败 - TASK_STATUS_CANCELED = 3; // 任务被取消 + WAIT = 0; // 任务等待执行 + RUNNING = 1; // 任务正在执行中 + SUCCESS = 2; // 任务成功完成 + FAILED = 3; // 任务执行失败 + CANCELLED = 4; // 任务被取消 } message CollectBinanceRequest { @@ -31,7 +32,7 @@ message CollectBinanceRequest { int32 chunk_size = 3; // 分块大小 } -message CollectBinaceByDateRequest { +message CollectBinanceByDateRequest { string task_id = 1; // 任务的唯一标识符 int32 start_ts = 2; // 起始时间 int32 end_ts = 3; // 终止时间 diff --git a/readme.md b/readme.md index 2961955..831489e 100644 --- a/readme.md +++ b/readme.md @@ -12,11 +12,13 @@ Etrade 是一个用于 CEX-DEX 套利分析的一站式平台,主要针对 Uniswap V3 和 Binance 的 ETH/USDT 交易对进行分析。 -当前版本的核心数据流/调用链路如下: +系统采用微服务架构,核心数据流如下: -`[Vue.js 前端 (浏览器)] -> [Go 后端 API (Gin)] -> [Python Worker (gRPC,采集/分析)] -> [PostgreSQL 数据库]` +`[Vue.js 前端] -> [Go 后端 API (Gin)] -> [Python Worker (gRPC)] -> [PostgreSQL 数据库]` -其中:Go 端负责任务创建/调度与对外 API;Python 端作为 Worker 执行具体脚本(采集/聚合/分析)并写入任务日志与结果。 +架构说明: +- **Go 后端**:负责任务创建、调度与对外 API 服务 +- **Python Worker**:执行数据采集、聚合、分析等具体任务,并写入任务日志与结果 ![](images/architecture.png) @@ -39,43 +41,45 @@ npm install # 下载对应依赖 npm run dev # 运行项目 ``` -前端默认请求后端 `http://localhost:8888/api/v1`(如需修改请使用前端环境变量配置)。 +前端默认请求后端 `http://localhost:8888/api/v1` -前端项目结构,主要是 src 目录: +**项目结构**(src 目录): -- api:接口,和后端对应,使用 `axios` 库 -- assets:一些公用的 css 等资源 -- components:可复用的组件 -- router:动态路由组件 -- views:vue页面 -- app.vue/main.tx/style.css:一些全局配置 -- views:vue 页面 -- app.vue/main.ts/style.css:一些全局配置 +- `api/`:API 接口层,使用 `axios` 与后端通信 +- `assets/`:公共静态资源(CSS 等) +- `components/`:可复用组件 +- `router/`:路由配置 +- `views/`:页面组件 +- `App.vue`、`main.ts`、`style.css`:应用入口和全局样式 -可能会在 `tsconfig.app.json` 这类配置文件里面出现一些很奇怪的报错,如果经检查确实没什么问题,很有可能是因为缓存机制,把报错的语句/文件删除了再恢复一般就正常了,实在有无法修复的奇怪报错可以忽略。 +> **提示**:如果 `tsconfig.app.json` 等配置文件中出现异常报错,经检查无实质性问题时,可能是缓存导致。可尝试删除并恢复相关语句/文件,通常即可解决。 -### postgresql +### PostgreSQL -可以使用docker拉取 -使用 docker 拉取,便于调整端口等配置。这里 postgresql 运行的端口用默认的 5432 端口(请确保这个端口可用,或者换到别的可用端口),默认用户名为 postgres,密码就是 123456,这个账号和密码用于访问数据库本身。 +使用 Docker 部署 PostgreSQL,便于配置管理: ```bash -# 拉取 PostgreSQL +# 拉取镜像 docker pull postgres -# 运行,配置尽量不要改 + +# 运行容器(默认端口 5432,用户名 postgres,密码 123456) docker run --name postgresql \ -e POSTGRES_PASSWORD=123456 \ -p 5432:5432 \ -d postgres ``` -PgAdmin(用于管理PostgreSQL)同样可以使用docker拉取,注意数据库的ip地址需要使用`host.docker.internal` -PgAdmin 用于查询和管理 PostgreSQL,同样可以使用 docker 拉取,下面的邮箱 `test@123.com` 和密码是用于访问 PgAdmin,但注意 PgAdmin 中输入 docker 部署的本地服务器 ip 地址时需要使用 `host.docker.internal`,而不是 `localhost` 或者 `127.0.0.1`。 +> **注意**:确保 5432 端口可用,或根据需要修改映射端口。 + +**PgAdmin(可选)** + +PgAdmin 用于可视化管理和查询 PostgreSQL,同样可通过 Docker 部署: ```bash -# 一并拉取 pgadmin4 方便查询 +# 拉取 PgAdmin 镜像 docker pull dpage/pgadmin4 +# 运行容器 docker run -d -p 5433:80 \ --name pgadmin4 \ -e PGADMIN_DEFAULT_EMAIL=test@123.com \ @@ -83,31 +87,55 @@ docker run -d -p 5433:80 \ dpage/pgadmin4 ``` +> **重要**:在 PgAdmin 中连接 Docker 部署的 PostgreSQL 时,主机地址应使用 `host.docker.internal`,而非 `localhost` 或 `127.0.0.1`。 + +### RabbitMQ + +使用 Docker 部署 RabbitMQ(含管理界面): + +```bash +docker run -d \ + --name rabbitmq \ + -p 5672:5672 \ + -p 15672:15672 \ + -e RABBITMQ_DEFAULT_USER=admin \ + -e RABBITMQ_DEFAULT_PASS=123456 \ + rabbitmq:management +``` + +管理界面地址:`http://localhost:15672`(用户名:admin,密码:123456) + ### 后端 -运行后端: +**启动步骤:** ```bash -# 在 backend 目录 -go mod tidy # 自动处理依赖关系 -go run main.go # 运行项目,端口 8888 +cd backend +go mod tidy # 安装依赖 +go run main.go # 启动服务(默认端口 8888) ``` -需要配置好`config/config.yaml`下的数据库连接信息 -并配置好: -- `backend/config/config.yaml`:数据库连接 + `worker.address`(Python Worker 地址) +**配置文件** + +需要配置 `backend/config/config.yaml`: +- 数据库连接信息 +- `worker.address`:Python Worker 地址 + +**API 文档** Swagger 地址:`http://localhost:8888/swagger/index.html` -项目结构,MVC 模式: -- api:用于收发 http 请求 -- db:数据库配置 -- models:模型,用于数据库存储和出入参 -- service:业务逻辑,可以与数据库交互/派发 Worker 任务 -- utils:一些工具方法 +**项目结构(MVC 模式)** -`utils/response.go`中定义了统一的后端返回方法,直接在api层调用这些方法返回即可,统一的格式为: -需要注意 `utils/response.go` 中定义了统一的后端返回方法,只需要直接在 api 层调用这些方法返回就行了,统一的格式为: +- `api/`:HTTP 请求处理层 +- `db/`:数据库连接配置 +- `models/`:数据模型(数据库映射与请求/响应结构) +- `service/`:业务逻辑层(数据库操作与 Worker 任务派发) +- `utils/`:工具函数 + +**统一响应格式** + +`utils/response.go` 定义了统一的 API 响应方法,响应格式如下: ```json { @@ -117,13 +145,16 @@ Swagger 地址:`http://localhost:8888/swagger/index.html` } ``` -比如要返回一个失败的请求,就可以: +使用示例: ```go +// 返回失败响应 utils.Fail(c, http.StatusInternalServerError, err.Error()) ``` -models 中的结构体注释建议都写,因为这个项目中理论上正常的数据每一项都是非空的,比如: +**数据模型规范** + +建议为所有模型字段添加 GORM 标签,特别是非空约束。例如: ```go type BinanceTrade struct { @@ -134,11 +165,11 @@ type BinanceTrade struct { } ``` -> 风险分析更新 (Risk Analysis Update) -> 我们在 `arbitrage_opportunities` 表中新增了 **`risk_metrics_json`** (JSONB) 字段。 -> 现在每次运行 `analyse` 任务时,除了计算利润,还会自动调用 Python 端的风险模型,计算包括 **滑点 (Slippage)**、**波动率 (Volatility)** 和 **风险评分 (Risk Score)** 等指标,并存入该字段。 +### 数据分析引擎 + +Python Worker 负责执行数据采集、聚合和分析任务。 -### data 分析引擎(Python Worker) +**启动步骤:** ```bash cd data @@ -146,7 +177,7 @@ pip install -r requirements.txt python server.py ``` -`data/server.py` 是 gRPC Worker:后端通过 `worker.address` 调用它执行采集/聚合/分析任务。 +`data/server.py` 作为 gRPC Worker 服务,后端通过 `worker.address` 调用其执行各类任务。 ## 系统部署指南