package giteaactions import ( "context" "net/http" "net/http/httptest" "reflect" "testing" "connectrpc.com/connect" "gitea.dev/actionslib/pkg/protocol" runnerv1 "gitea.dev/actionslib/runner/v1" "gitea.dev/actionslib/runner/v1/runnerv1connect" ) type runnerService struct { runnerv1connect.UnimplementedRunnerServiceHandler t *testing.T declaredLabels []string fetchedVersion int64 updatedTask int64 updatedLog int64 expectedUUID string expectedToken string } func (s *runnerService) checkAuth(request connect.AnyRequest) { s.t.Helper() if got := request.Header().Get(protocol.UUIDHeader); got != s.expectedUUID { s.t.Errorf("runner UUID header = %q", got) } if got := request.Header().Get(protocol.TokenHeader); got != s.expectedToken { s.t.Errorf("runner token header = %q", got) } } func (s *runnerService) Declare(_ context.Context, request *connect.Request[runnerv1.DeclareRequest]) (*connect.Response[runnerv1.DeclareResponse], error) { s.checkAuth(request) s.declaredLabels = request.Msg.Labels return connect.NewResponse(&runnerv1.DeclareResponse{}), nil } func (s *runnerService) FetchTask(_ context.Context, request *connect.Request[runnerv1.FetchTaskRequest]) (*connect.Response[runnerv1.FetchTaskResponse], error) { s.checkAuth(request) s.fetchedVersion = request.Msg.TasksVersion return connect.NewResponse(&runnerv1.FetchTaskResponse{ Task: &runnerv1.Task{Id: 42}, TasksVersion: 8, }), nil } func (s *runnerService) UpdateTask(_ context.Context, request *connect.Request[runnerv1.UpdateTaskRequest]) (*connect.Response[runnerv1.UpdateTaskResponse], error) { s.checkAuth(request) s.updatedTask = request.Msg.GetState().GetId() return connect.NewResponse(&runnerv1.UpdateTaskResponse{State: request.Msg.State}), nil } func (s *runnerService) UpdateLog(_ context.Context, request *connect.Request[runnerv1.UpdateLogRequest]) (*connect.Response[runnerv1.UpdateLogResponse], error) { s.checkAuth(request) s.updatedLog = request.Msg.GetTaskId() return connect.NewResponse(&runnerv1.UpdateLogResponse{}), nil } func TestClientUsesOfficialRunnerProtocol(t *testing.T) { service := &runnerService{ t: t, expectedUUID: "runner-uuid", expectedToken: "runner-token", } path, handler := runnerv1connect.NewRunnerServiceHandler(service) mux := http.NewServeMux() mux.Handle("/api/actions"+path, http.StripPrefix("/api/actions", handler)) server := httptest.NewServer(mux) defer server.Close() client := NewClient(server.Client(), server.URL, service.expectedUUID, service.expectedToken) labels := []string{"self-hosted", "pod", "vm"} if err := client.Declare(context.Background(), "0.1.0", labels); err != nil { t.Fatal(err) } response, err := client.FetchTask(context.Background(), 7) if err != nil { t.Fatal(err) } if !reflect.DeepEqual(service.declaredLabels, labels) { t.Fatalf("declared labels = %#v", service.declaredLabels) } if service.fetchedVersion != 7 || response.GetTasksVersion() != 8 || response.GetTask().GetId() != 42 { t.Fatalf("unexpected FetchTask exchange: request=%d response=%v", service.fetchedVersion, response) } if _, err := client.UpdateTask(context.Background(), connect.NewRequest(&runnerv1.UpdateTaskRequest{State: &runnerv1.TaskState{Id: 42}})); err != nil { t.Fatal(err) } if _, err := client.UpdateLog(context.Background(), connect.NewRequest(&runnerv1.UpdateLogRequest{TaskId: 42})); err != nil { t.Fatal(err) } if service.updatedTask != 42 || service.updatedLog != 42 { t.Fatalf("updated task=%d log=%d", service.updatedTask, service.updatedLog) } }