Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
112 changes: 110 additions & 2 deletions backend/pkg/controller/flows.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package controller

import (
"context"
"database/sql"
"errors"
"fmt"
"sort"
Expand Down Expand Up @@ -348,7 +349,7 @@ func (fc *flowController) StopFlow(ctx context.Context, flowID int64) error {

flow, ok := fc.flows[flowID]
if !ok {
return ErrFlowNotFound
return fc.stopUnloadedFlow(ctx, flowID)
}

err := flow.Stop(ctx)
Expand All @@ -365,7 +366,7 @@ func (fc *flowController) FinishFlow(ctx context.Context, flowID int64) error {

flow, ok := fc.flows[flowID]
if !ok {
return ErrFlowNotFound
return fc.finishUnloadedFlow(ctx, flowID)
}

err := flow.Finish(ctx)
Expand All @@ -378,6 +379,113 @@ func (fc *flowController) FinishFlow(ctx context.Context, flowID int64) error {
return nil
}

func (fc *flowController) stopUnloadedFlow(ctx context.Context, flowID int64) error {
flow, err := fc.db.GetFlow(ctx, flowID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return ErrFlowNotFound
}
return fmt.Errorf("failed to get flow %d: %w", flowID, err)
}

containers, err := fc.cleanupFlowContainers(ctx, flowID, func(status database.ContainerStatus) bool {
return status == database.ContainerStatusStarting || status == database.ContainerStatusRunning
})
if err != nil {
return err
}

switch flow.Status {
case database.FlowStatusRunning, database.FlowStatusWaiting:
flow, err = fc.db.UpdateFlowStatus(ctx, database.UpdateFlowStatusParams{
ID: flowID,
Status: database.FlowStatusWaiting,
})
if err != nil {
return fmt.Errorf("failed to set flow %d status to waiting: %w", flowID, err)
}
fc.publishFlowUpdated(ctx, flow, containers)
return nil
default:
fc.publishFlowUpdated(ctx, flow, containers)
return nil
}
}

func (fc *flowController) finishUnloadedFlow(ctx context.Context, flowID int64) error {
if _, err := fc.db.GetFlow(ctx, flowID); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return ErrFlowNotFound
}
return fmt.Errorf("failed to get flow %d: %w", flowID, err)
}

containers, err := fc.cleanupFlowContainers(ctx, flowID, func(status database.ContainerStatus) bool {
return status != database.ContainerStatusDeleted
})
if err != nil {
return err
}

flow, err := fc.db.UpdateFlowStatus(ctx, database.UpdateFlowStatusParams{
ID: flowID,
Status: database.FlowStatusFinished,
})
if err != nil {
return fmt.Errorf("failed to set flow %d status to finished: %w", flowID, err)
}

fc.publishFlowUpdated(ctx, flow, containers)

return nil
}

func (fc *flowController) cleanupFlowContainers(
ctx context.Context,
flowID int64,
shouldCleanup func(database.ContainerStatus) bool,
) ([]database.Container, error) {
containers, err := fc.db.GetFlowContainers(ctx, flowID)
if err != nil {
return nil, fmt.Errorf("failed to get flow %d containers: %w", flowID, err)
}

for _, container := range containers {
if !shouldCleanup(container.Status) {
continue
}

if container.LocalID.Valid && container.LocalID.String != "" && fc.docker != nil {
if err := fc.docker.RemoveContainer(ctx, container.LocalID.String, container.ID); err != nil {
return nil, fmt.Errorf("failed to remove flow %d container %d: %w", flowID, container.ID, err)
}
continue
}

if _, err := fc.db.UpdateContainerStatus(ctx, database.UpdateContainerStatusParams{
ID: container.ID,
Status: database.ContainerStatusDeleted,
}); err != nil {
return nil, fmt.Errorf("failed to mark flow %d container %d as deleted: %w", flowID, container.ID, err)
}
}

containers, err = fc.db.GetFlowContainers(ctx, flowID)
if err != nil {
return nil, fmt.Errorf("failed to get flow %d containers after cleanup: %w", flowID, err)
}

return containers, nil
}

func (fc *flowController) publishFlowUpdated(ctx context.Context, flow database.Flow, containers []database.Container) {
if fc.subs == nil {
return
}

fc.subs.NewFlowPublisher(flow.UserID, flow.ID).FlowUpdated(ctx, flow, containers)
}

func (fc *flowController) RenameFlow(ctx context.Context, flowID int64, title string) error {
fc.mx.Lock()
defer fc.mx.Unlock()
Expand Down
223 changes: 223 additions & 0 deletions backend/pkg/controller/flows_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,223 @@
package controller

import (
"context"
"database/sql"
"errors"
"reflect"
"sync"
"testing"

"pentagi/pkg/database"
"pentagi/pkg/docker"
)

type fallbackQuerier struct {
database.Querier

flow database.Flow
flowExists bool
containers []database.Container
}

func (q *fallbackQuerier) GetFlow(ctx context.Context, id int64) (database.Flow, error) {
if !q.flowExists || q.flow.ID != id {
return database.Flow{}, sql.ErrNoRows
}

return q.flow, nil
}

func (q *fallbackQuerier) GetFlowContainers(ctx context.Context, flowID int64) ([]database.Container, error) {
containers := make([]database.Container, 0, len(q.containers))
for _, container := range q.containers {
if container.FlowID == flowID {
containers = append(containers, container)
}
}

return containers, nil
}

func (q *fallbackQuerier) UpdateFlowStatus(ctx context.Context, arg database.UpdateFlowStatusParams) (database.Flow, error) {
if !q.flowExists || q.flow.ID != arg.ID {
return database.Flow{}, sql.ErrNoRows
}

q.flow.Status = arg.Status

return q.flow, nil
}

func (q *fallbackQuerier) UpdateContainerStatus(
ctx context.Context,
arg database.UpdateContainerStatusParams,
) (database.Container, error) {
for idx := range q.containers {
if q.containers[idx].ID == arg.ID {
q.containers[idx].Status = arg.Status
return q.containers[idx], nil
}
}

return database.Container{}, sql.ErrNoRows
}

type fallbackDockerClient struct {
docker.DockerClient

db *fallbackQuerier
removed []string
}

func (c *fallbackDockerClient) RemoveContainer(ctx context.Context, containerID string, dbID int64) error {
c.removed = append(c.removed, containerID)
_, err := c.db.UpdateContainerStatus(ctx, database.UpdateContainerStatusParams{
ID: dbID,
Status: database.ContainerStatusDeleted,
})

return err
}

func TestFlowControllerFinishFlowFallbackRemovesContainerAndMarksFinished(t *testing.T) {
ctx := context.Background()
db := &fallbackQuerier{
flowExists: true,
flow: database.Flow{
ID: 42,
UserID: 7,
Status: database.FlowStatusRunning,
},
containers: []database.Container{
{
ID: 11,
FlowID: 42,
Status: database.ContainerStatusRunning,
LocalID: database.StringToNullString("docker-11"),
},
},
}
dockerClient := &fallbackDockerClient{db: db}
fc := &flowController{
db: db,
mx: &sync.Mutex{},
flows: map[int64]FlowWorker{},
docker: dockerClient,
}

if err := fc.FinishFlow(ctx, 42); err != nil {
t.Fatalf("FinishFlow returned error: %v", err)
}

if db.flow.Status != database.FlowStatusFinished {
t.Fatalf("expected flow status %q, got %q", database.FlowStatusFinished, db.flow.Status)
}
if db.containers[0].Status != database.ContainerStatusDeleted {
t.Fatalf("expected container status %q, got %q", database.ContainerStatusDeleted, db.containers[0].Status)
}
if !reflect.DeepEqual(dockerClient.removed, []string{"docker-11"}) {
t.Fatalf("expected docker removal for docker-11, got %#v", dockerClient.removed)
}
}

func TestFlowControllerFinishFlowFallbackMissingFlow(t *testing.T) {
ctx := context.Background()
db := &fallbackQuerier{}
fc := &flowController{
db: db,
mx: &sync.Mutex{},
flows: map[int64]FlowWorker{},
}

if err := fc.FinishFlow(ctx, 404); !errors.Is(err, ErrFlowNotFound) {
t.Fatalf("expected ErrFlowNotFound, got %v", err)
}
}

func TestFlowControllerStopFlowFallbackRemovesContainerAndMarksWaiting(t *testing.T) {
ctx := context.Background()
db := &fallbackQuerier{
flowExists: true,
flow: database.Flow{
ID: 42,
UserID: 7,
Status: database.FlowStatusRunning,
},
containers: []database.Container{
{
ID: 11,
FlowID: 42,
Status: database.ContainerStatusRunning,
LocalID: database.StringToNullString("docker-11"),
},
{
ID: 12,
FlowID: 42,
Status: database.ContainerStatusStopped,
LocalID: database.StringToNullString("docker-12"),
},
},
}
dockerClient := &fallbackDockerClient{db: db}
fc := &flowController{
db: db,
mx: &sync.Mutex{},
flows: map[int64]FlowWorker{},
docker: dockerClient,
}

if err := fc.StopFlow(ctx, 42); err != nil {
t.Fatalf("StopFlow returned error: %v", err)
}

if db.flow.Status != database.FlowStatusWaiting {
t.Fatalf("expected flow status %q, got %q", database.FlowStatusWaiting, db.flow.Status)
}
if db.containers[0].Status != database.ContainerStatusDeleted {
t.Fatalf("expected running container status %q, got %q", database.ContainerStatusDeleted, db.containers[0].Status)
}
if db.containers[1].Status != database.ContainerStatusStopped {
t.Fatalf("expected stopped container status %q, got %q", database.ContainerStatusStopped, db.containers[1].Status)
}
if !reflect.DeepEqual(dockerClient.removed, []string{"docker-11"}) {
t.Fatalf("expected docker removal for docker-11, got %#v", dockerClient.removed)
}
}

func TestFlowControllerFallbackMarksContainerWithoutLocalIDDeleted(t *testing.T) {
ctx := context.Background()
db := &fallbackQuerier{
flowExists: true,
flow: database.Flow{
ID: 42,
UserID: 7,
Status: database.FlowStatusRunning,
},
containers: []database.Container{
{
ID: 11,
FlowID: 42,
Status: database.ContainerStatusRunning,
},
},
}
dockerClient := &fallbackDockerClient{db: db}
fc := &flowController{
db: db,
mx: &sync.Mutex{},
flows: map[int64]FlowWorker{},
docker: dockerClient,
}

if err := fc.FinishFlow(ctx, 42); err != nil {
t.Fatalf("FinishFlow returned error: %v", err)
}

if db.containers[0].Status != database.ContainerStatusDeleted {
t.Fatalf("expected container status %q, got %q", database.ContainerStatusDeleted, db.containers[0].Status)
}
if len(dockerClient.removed) != 0 {
t.Fatalf("expected no docker removals, got %#v", dockerClient.removed)
}
}
6 changes: 1 addition & 5 deletions backend/pkg/graph/schema.resolvers.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading
Loading