package controller

import (
	"context"
	"database/sql"
	"errors"
	"fmt"
	"sync"
	"time"

	"pentagi/pkg/database"
	obs "pentagi/pkg/observability"
	"pentagi/pkg/providers"
	"pentagi/pkg/tools"

	"github.com/sirupsen/logrus"
)

type FlowUpdater interface {
	SetStatus(ctx context.Context, status database.FlowStatus) error
}

type TaskWorker interface {
	GetTaskID() int64
	GetFlowID() int64
	GetUserID() int64
	GetTitle() string
	IsCompleted() bool
	IsWaiting() bool
	GetStatus(ctx context.Context) (database.TaskStatus, error)
	SetStatus(ctx context.Context, status database.TaskStatus) error
	GetResult(ctx context.Context) (string, error)
	Report(ctx context.Context) error
	SetResult(ctx context.Context, result string) error
	PutInput(ctx context.Context, input string) error
	Run(ctx context.Context) error
	Finish(ctx context.Context) error
	InvalidateSubtasks(subtaskIDs []int64)
}

type taskWorker struct {
	mx        *sync.RWMutex
	stc       SubtaskController
	taskCtx   *TaskContext
	updater   FlowUpdater
	completed bool
	waiting   bool
}

func NewTaskWorker(
	ctx context.Context,
	flowCtx *FlowContext,
	input string,
	updater FlowUpdater,
) (_ TaskWorker, err error) {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.NewTaskWorker")
	defer span.End()

	ctx = tools.PutAgentContext(ctx, database.MsgchainTypePrimaryAgent)

	title, err := flowCtx.Provider.GetTaskTitle(ctx, input)
	if err != nil {
		return nil, fmt.Errorf("failed to get task title: %w", err)
	}

	task, err := flowCtx.DB.CreateTask(ctx, database.CreateTaskParams{
		Status: database.TaskStatusCreated,
		Title:  title,
		Input:  input,
		FlowID: flowCtx.FlowID,
	})
	if err != nil {
		return nil, fmt.Errorf("failed to create task in DB: %w", err)
	}
	defer func() {
		if err != nil {
			failCreatedTask(ctx, flowCtx, task, err)
		}
	}()

	flowCtx.Publisher.TaskCreated(ctx, task, []database.Subtask{})

	taskCtx := &TaskContext{
		FlowContext: *flowCtx,
		TaskID:      task.ID,
		TaskTitle:   title,
		TaskInput:   input,
	}
	stc := NewSubtaskController(taskCtx)

	_, err = taskCtx.MsgLog.PutTaskMsg(
		ctx,
		database.MsglogTypeInput,
		taskCtx.TaskID,
		"", // thinking is empty because this is input
		input,
	)
	if err != nil {
		return nil, fmt.Errorf("failed to put input for task %d: %w", taskCtx.TaskID, err)
	}

	err = stc.GenerateSubtasks(ctx)
	if err != nil {
		return nil, fmt.Errorf("failed to generate subtasks: %w", err)
	}

	subtasks, err := flowCtx.DB.GetTaskSubtasks(ctx, task.ID)
	if err != nil {
		return nil, fmt.Errorf("failed to get subtasks for task %d: %w", task.ID, err)
	}

	flowCtx.Publisher.TaskUpdated(ctx, task, subtasks)

	return &taskWorker{
		mx:        &sync.RWMutex{},
		stc:       stc,
		taskCtx:   taskCtx,
		updater:   updater,
		completed: false,
		waiting:   false,
	}, nil
}

func failCreatedTask(ctx context.Context, flowCtx *FlowContext, task database.Task, cause error) {
	ctx = context.WithoutCancel(ctx)

	var (
		closed database.Task
		err    error
	)
	if errors.Is(cause, context.Canceled) {
		closed, err = flowCtx.DB.UpdateTaskStatus(ctx, database.UpdateTaskStatusParams{
			Status: database.TaskStatusFinished,
			ID:     task.ID,
		})
	} else {
		closed, err = flowCtx.DB.UpdateTaskFailedResult(ctx, database.UpdateTaskFailedResultParams{
			Result: database.SanitizeUTF8(cause.Error()),
			ID:     task.ID,
		})
	}
	if err != nil {
		logrus.WithContext(ctx).WithError(err).
			Warnf("failed to close task %d after it could not start", task.ID)
		return
	}

	subtasks, err := flowCtx.DB.GetTaskSubtasks(ctx, task.ID)
	if err != nil {
		logrus.WithContext(ctx).WithError(err).
			Warnf("failed to get subtasks of task %d, publishing it closed without them", task.ID)
	}

	flowCtx.Publisher.TaskUpdated(ctx, closed, subtasks)
}

func LoadTaskWorker(
	ctx context.Context,
	task database.Task,
	flowCtx *FlowContext,
	updater FlowUpdater,
) (TaskWorker, error) {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.LoadTaskWorker")
	defer span.End()

	ctx = tools.PutAgentContext(ctx, database.MsgchainTypePrimaryAgent)
	taskCtx := &TaskContext{
		FlowContext: *flowCtx,
		TaskID:      task.ID,
		TaskTitle:   task.Title,
		TaskInput:   task.Input,
	}

	stc := NewSubtaskController(taskCtx)
	var completed, waiting bool
	switch task.Status {
	case database.TaskStatusFinished, database.TaskStatusFailed:
		completed = true
	case database.TaskStatusWaiting:
		waiting = true
	case database.TaskStatusRunning:
	case database.TaskStatusCreated:
		return nil, fmt.Errorf("task %d has created yet: loading aborted: %w", task.ID, ErrNothingToLoad)
	}

	tw := &taskWorker{
		mx:        &sync.RWMutex{},
		stc:       stc,
		taskCtx:   taskCtx,
		updater:   updater,
		completed: completed,
		waiting:   waiting,
	}

	if err := tw.stc.LoadSubtasks(ctx, task.ID, tw); err != nil {
		return nil, fmt.Errorf("failed to load subtasks for task %d: %w", task.ID, err)
	}

	return tw, nil
}

func (tw *taskWorker) GetTaskID() int64 {
	return tw.taskCtx.TaskID
}

func (tw *taskWorker) GetFlowID() int64 {
	return tw.taskCtx.FlowID
}

func (tw *taskWorker) GetUserID() int64 {
	return tw.taskCtx.UserID
}

func (tw *taskWorker) GetTitle() string {
	return tw.taskCtx.TaskTitle
}

func (tw *taskWorker) IsCompleted() bool {
	tw.mx.RLock()
	defer tw.mx.RUnlock()

	return tw.completed
}

func (tw *taskWorker) IsWaiting() bool {
	tw.mx.RLock()
	defer tw.mx.RUnlock()

	return tw.waiting
}

func (tw *taskWorker) GetStatus(ctx context.Context) (database.TaskStatus, error) {
	task, err := tw.taskCtx.DB.GetTask(ctx, tw.taskCtx.TaskID)
	if err != nil {
		return database.TaskStatusFailed, err
	}

	return task.Status, nil
}

// this function is exclusively change task internal properties "completed" and "waiting"
func (tw *taskWorker) SetStatus(ctx context.Context, status database.TaskStatus) error {
	task, err := tw.taskCtx.DB.UpdateTaskStatus(ctx, database.UpdateTaskStatusParams{
		Status: status,
		ID:     tw.taskCtx.TaskID,
	})
	if err != nil {
		if errors.Is(err, sql.ErrNoRows) {
			// Replacement can leave this worker behind after deleting its row.
			// Treat it as complete so flow shutdown remains idempotent.
			logrus.WithContext(ctx).WithFields(logrus.Fields{
				"task_id":          tw.taskCtx.TaskID,
				"requested_status": status,
			}).Warn("task no longer exists in the database, treating as already finished")

			tw.mx.Lock()
			tw.completed = true
			tw.waiting = false
			tw.mx.Unlock()

			return nil
		}

		return fmt.Errorf("failed to set task %d status: %w", tw.taskCtx.TaskID, err)
	}

	subtasks, err := tw.taskCtx.DB.GetTaskSubtasks(ctx, tw.taskCtx.TaskID)
	if err != nil {
		return fmt.Errorf("failed to get task %d subtasks: %w", tw.taskCtx.TaskID, err)
	}

	tw.taskCtx.Publisher.TaskUpdated(ctx, task, subtasks)

	tw.mx.Lock()
	defer tw.mx.Unlock()

	switch status {
	case database.TaskStatusRunning:
		tw.completed = false
		tw.waiting = false
		err = tw.updater.SetStatus(ctx, database.FlowStatusRunning)
	case database.TaskStatusWaiting:
		tw.completed = false
		tw.waiting = true
		err = tw.updater.SetStatus(ctx, database.FlowStatusWaiting)
	case database.TaskStatusFinished, database.TaskStatusFailed:
		tw.completed = true
		tw.waiting = false
		// the last task was done, set flow status to Waiting new user input
		err = tw.updater.SetStatus(ctx, database.FlowStatusWaiting)
	default:
		// status Created is not possible to set by this call
		return fmt.Errorf("unsupported task status: %s", status)
	}
	if err != nil {
		return fmt.Errorf("failed to set flow status in back propagation: %w", err)
	}

	return nil
}

func (tw *taskWorker) GetResult(ctx context.Context) (string, error) {
	task, err := tw.taskCtx.DB.GetTask(ctx, tw.taskCtx.TaskID)
	if err != nil {
		return "", err
	}

	return task.Result, nil
}

func (tw *taskWorker) SetResult(ctx context.Context, result string) error {
	_, err := tw.taskCtx.DB.UpdateTaskResult(ctx, database.UpdateTaskResultParams{
		Result: result,
		ID:     tw.taskCtx.TaskID,
	})
	if err != nil {
		return fmt.Errorf("failed to set task %d result: %w", tw.taskCtx.TaskID, err)
	}

	return nil
}

func (tw *taskWorker) PutInput(ctx context.Context, input string) error {
	if !tw.IsWaiting() {
		return fmt.Errorf("task is not waiting")
	}

	for _, st := range tw.stc.ListSubtasks(ctx) {
		if !st.IsCompleted() && st.IsWaiting() {
			if err := st.PutInput(ctx, input); err != nil {
				return fmt.Errorf("failed to put input to subtask %d: %w", st.GetSubtaskID(), err)
			} else {
				break
			}
		}
	}

	return nil
}

func (tw *taskWorker) Run(ctx context.Context) error {
	ctx = tools.PutAgentContext(ctx, database.MsgchainTypePrimaryAgent)

	for len(tw.stc.ListSubtasks(ctx)) < providers.TasksNumberLimit+3 {
		st, err := tw.stc.PopSubtask(ctx, tw)
		if err != nil {
			tw.handleInterrupting(err)
			return err
		}

		// empty queue for subtasks means that task is done
		if st == nil {
			break
		}

		if err := st.Run(ctx); err != nil {
			tw.handleInterrupting(err)
			return err
		}

		// pass through if task is waiting from back status propagation
		if tw.IsWaiting() {
			return nil
		} // otherwise subtask is done

		if err := tw.stc.RefineSubtasks(ctx); err != nil {
			if errors.Is(err, context.Canceled) {
				ctx = context.Background()
			}
			_ = tw.SetStatus(ctx, database.TaskStatusWaiting)
			return fmt.Errorf("failed to refine subtasks list for the task %d: %w", tw.taskCtx.TaskID, err)
		}
	}

	return tw.writeResult(ctx)
}

// writeResult asks the reporter for the task's write-up and publishes it: the
// result, the task status it implies, and the report message the operator
// reads. Shared with Report, which is the same work asked for on demand
// instead of at the end of the subtask queue.
func (tw *taskWorker) writeResult(ctx context.Context) error {
	jobResult, err := tw.taskCtx.Provider.GetTaskResult(ctx, tw.taskCtx.TaskID)
	if err != nil {
		tw.handleInterrupting(err)
		return fmt.Errorf("failed to get task %d result: %w", tw.taskCtx.TaskID, err)
	}

	var taskStatus database.TaskStatus
	if jobResult.Success.Bool() {
		taskStatus = database.TaskStatusFinished
	} else {
		taskStatus = database.TaskStatusFailed
	}

	if err := tw.SetResult(ctx, jobResult.Result); err != nil {
		tw.handleInterrupting(err)
		return err
	}

	if err := tw.SetStatus(ctx, taskStatus); err != nil {
		tw.handleInterrupting(err)
		return err
	}

	format := database.MsglogResultFormatMarkdown
	_, err = tw.taskCtx.MsgLog.PutTaskMsgResult(
		ctx,
		database.MsglogTypeReport,
		tw.taskCtx.TaskID,
		"", // thinking is empty because agent can't return it
		tw.taskCtx.TaskTitle,
		jobResult.Result,
		format,
	)
	if err != nil {
		tw.handleInterrupting(err)
		return fmt.Errorf("failed to put report for task %d: %w", tw.taskCtx.TaskID, err)
	}

	return nil
}

// ErrTaskAlreadyCompleted is a report asked for on a task that has nothing left
// to write up. It is the caller's mistake rather than a fault, and the retry
// loop in flowWorker.runReport stops on it rather than asking three times.
var ErrTaskAlreadyCompleted = errors.New("task has already completed")

// Report writes the task up on demand, for a task still open -- running or
// waiting alike. The write-up is the one artefact a flow produces that cannot
// be re-derived once the flow is finished, so this exists to ask for it before
// that point.
//
// The caller bounds the run: the context it passes is the whole allowance,
// because a retrying caller must not multiply it.
func (tw *taskWorker) Report(ctx context.Context) error {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.taskWorker.Report")
	defer span.End()

	if tw.IsCompleted() {
		return fmt.Errorf("%w: task %d", ErrTaskAlreadyCompleted, tw.taskCtx.TaskID)
	}

	// Run's first statement, so a report-initiated chain is attributed the same
	// way the ordinary one is.
	ctx = tools.PutAgentContext(ctx, database.MsgchainTypePrimaryAgent)

	if err := tw.SetStatus(ctx, database.TaskStatusRunning); err != nil {
		return fmt.Errorf("failed to start reporting task %d: %w", tw.taskCtx.TaskID, err)
	}

	if err := tw.writeResult(ctx); err != nil {
		// Only when the write-up did not already close the task. writeResult
		// commits the status before the report message, so a failure on that
		// last step leaves a task that IS finished -- and resetting it here
		// would un-finish it and buy a second reporter run. This is the same
		// invariant handleInterrupting protects.
		if !tw.IsCompleted() {
			resetCtx := ctx
			if ctx.Err() != nil {
				resetCtx = context.WithoutCancel(ctx)
			}
			_ = tw.SetStatus(resetCtx, database.TaskStatusWaiting)
		}

		return err
	}

	// A task closed by a report has the same leftovers an ordinary finish
	// closes: writeResult marks the task done without touching subtasks that
	// never ran, and they would sit under a finished task forever.
	return tw.finishOutstandingSubtasks(ctx)
}

// finishOutstandingSubtasks closes any subtask the task did not reach. Shared
// wording with Finish, which does the same before marking the task done.
func (tw *taskWorker) finishOutstandingSubtasks(ctx context.Context) error {
	for _, st := range tw.stc.ListSubtasks(ctx) {
		if st.IsCompleted() {
			continue
		}
		if err := st.Finish(ctx); err != nil {
			return fmt.Errorf("failed to finish subtask %d after reporting: %w", st.GetSubtaskID(), err)
		}
	}

	return nil
}

// handleInterrupting sets this task (and the flow via taskWorker.SetStatus)
// to Waiting when err is context.Canceled or context.DeadlineExceeded. Skips if the task is
// already marked completed in memory (Finished/Failed) so we do not revive a finished task.
func (tw *taskWorker) handleInterrupting(err error) {
	if err == nil {
		return
	}
	if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
		return
	}
	if tw.IsCompleted() {
		return
	}

	resetCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
	defer cancel()

	if errSt := tw.SetStatus(resetCtx, database.TaskStatusWaiting); errSt != nil {
		logrus.WithError(errSt).Warn("failed to set task waiting after run interrupt")
	}
}

func (tw *taskWorker) Finish(ctx context.Context) error {
	if tw.IsCompleted() {
		return fmt.Errorf("task has already completed")
	}

	for _, st := range tw.stc.ListSubtasks(ctx) {
		if !st.IsCompleted() {
			if err := st.Finish(ctx); err != nil {
				return err
			}
		}
	}

	if err := tw.SetStatus(ctx, database.TaskStatusFinished); err != nil {
		return err
	}

	return nil
}

// InvalidateSubtasks evicts subtaskIDs from this task's in-memory subtask
// controller. See flowWorker.InvalidateTaskSubtasks for why this is needed.
func (tw *taskWorker) InvalidateSubtasks(subtaskIDs []int64) {
	tw.stc.InvalidateSubtasks(subtaskIDs)
}
