package controller

import (
	"context"
	"database/sql"
	"encoding/json"
	"errors"
	"fmt"
	"sync"
	"sync/atomic"
	"time"

	"pentagi/pkg/cast"
	"pentagi/pkg/database"
	"pentagi/pkg/graph/subscriptions"
	obs "pentagi/pkg/observability"
	"pentagi/pkg/observability/langfuse"
	"pentagi/pkg/providers"
	"pentagi/pkg/providers/pconfig"
	"pentagi/pkg/providers/provider"
	"pentagi/pkg/templates"
	"pentagi/pkg/tools"

	"github.com/sirupsen/logrus"
)

const stopAssistantTimeout = 5 * time.Second

type AssistantWorker interface {
	GetAssistantID() int64
	GetUserID() int64
	GetFlowID() int64
	GetTitle() string
	GetStatus(ctx context.Context) (database.AssistantStatus, error)
	SetStatus(ctx context.Context, status database.AssistantStatus) error
	PutInput(ctx context.Context, input string, useAgents bool, resources []database.UserResource) error
	Finish(ctx context.Context) error
	Stop(ctx context.Context) error
}

type assistantWorker struct {
	id      int64
	flowID  int64
	userID  int64
	chainID int64
	aslw    FlowAssistantLogWorker
	ap      providers.AssistantProvider
	db      database.Querier
	wg      *sync.WaitGroup
	pub     subscriptions.FlowPublisher
	fw      FlowWorker
	ctx     context.Context
	cancel  context.CancelFunc
	runMX   *sync.Mutex
	runST   context.CancelFunc
	runWG   *sync.WaitGroup
	input   chan assistantInput
	logger  *logrus.Entry

	noteOwed atomic.Bool
}

type newAssistantWorkerCtx struct {
	userID    int64
	flowID    int64
	input     string
	useAgents bool
	prvname   provider.ProviderName
	prvtype   provider.ProviderType
	functions *tools.Functions
	resources []database.UserResource
	fw        FlowWorker

	flowWorkerCtx
}

type assistantWorkerCtx struct {
	userID int64
	flowID int64
	fw     FlowWorker

	// prompter is an optional reusable prompter for the user. When non-nil,
	// LoadAssistantWorker uses it instead of issuing another GetUserPrompts
	// query and merging defaults. LoadFlowWorker sets this so a flow load
	// with multiple assistants only pays the DB+merge cost once.
	prompter templates.Prompter

	flowWorkerCtx
}

const assistantInputTimeout = 2 * time.Second

const interruptedRunNote = "The server restarted before the assistant finished this request, " +
	"so its reply may be incomplete. Send a message to continue."

type assistantInput struct {
	input     string
	useAgents bool
	reply     *inputReply
}

func reserveAssistant(ctx context.Context, awc newAssistantWorkerCtx) (database.Assistant, error) {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.reserveAssistant")
	defer span.End()

	assistant, err := awc.db.CreateAssistant(ctx, database.CreateAssistantParams{
		Title:              "untitled",
		Status:             database.AssistantStatusCreated,
		Model:              "unknown",
		ModelProviderName:  string(awc.prvname),
		ModelProviderType:  database.ProviderType(awc.prvtype),
		Language:           "English",
		ToolCallIDTemplate: cast.ToolCallIDTemplate,
		Functions:          []byte("{}"),
		FlowID:             awc.flowID,
		UseAgents:          awc.useAgents,
	})
	if err != nil {
		logrus.WithContext(ctx).WithError(err).Error("failed to create assistant in DB")
		return database.Assistant{}, fmt.Errorf("failed to create assistant in DB: %w", err)
	}

	logrus.WithContext(ctx).WithFields(logrus.Fields{
		"flow_id":      awc.flowID,
		"user_id":      awc.userID,
		"assistant_id": assistant.ID,
	}).Info("assistant created in DB")

	return assistant, nil
}

func buildAssistantWorker(
	ctx context.Context, assistant database.Assistant, awc newAssistantWorkerCtx,
) (_ AssistantWorker, retErr error) {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.buildAssistantWorker")
	defer span.End()

	logger := logrus.WithContext(ctx).WithFields(logrus.Fields{
		"flow_id":       awc.flowID,
		"user_id":       awc.userID,
		"assistant_id":  assistant.ID,
		"provider_name": awc.prvname.String(),
		"provider_type": awc.prvtype.String(),
	})

	user, err := awc.db.GetUser(ctx, awc.userID)
	if err != nil {
		logger.WithError(err).Error("failed to get user")
		return nil, fmt.Errorf("failed to get user %d: %w", awc.userID, err)
	}

	container, err := awc.db.GetFlowPrimaryContainer(ctx, awc.flowID)
	if err != nil {
		logger.WithError(err).Error("failed to get flow primary container")
		return nil, fmt.Errorf("failed to get flow primary container: %w", err)
	}

	pub := awc.subs.NewFlowPublisher(awc.userID, awc.flowID)

	defer func() {
		if retErr == nil || errors.Is(retErr, ErrFlowAlreadyStopped) {
			return
		}

		// The worker's own context may already be cancelled by the cleanup below.
		cleanupCtx := context.WithoutCancel(ctx)

		failed, statusErr := awc.db.UpdateAssistantStatus(cleanupCtx, database.UpdateAssistantStatusParams{
			ID:     assistant.ID,
			Status: database.AssistantStatusFailed,
		})
		if statusErr != nil {
			if !errors.Is(statusErr, sql.ErrNoRows) {
				logger.WithError(statusErr).Error("failed to mark the unstarted assistant failed")
			}

			return
		}

		pub.AssistantUpdated(cleanupCtx, failed)
	}()

	ctx, observation := obs.Observer.NewObservation(ctx,
		langfuse.WithObservationTraceContext(
			langfuse.WithTraceName(fmt.Sprintf("%s%d flow %d assistant worker", awc.cfg.TenantLabel(), awc.flowID, assistant.ID)),
			langfuse.WithTraceUserID(tenantUserID(awc.cfg, user.Mail)),
			langfuse.WithTraceTags(tenantTags(awc.cfg, "controller", "assistant")),
			langfuse.WithTraceInput(awc.input),
			langfuse.WithTraceSessionID(awc.cfg.ScopedName(fmt.Sprintf("assistant-%d-flow-%d", assistant.ID, awc.flowID))),
			langfuse.WithTraceMetadata(tenantMeta(awc.cfg, langfuse.Metadata{
				"assistant_id":  assistant.ID,
				"flow_id":       awc.flowID,
				"user_id":       awc.userID,
				"user_email":    user.Mail,
				"user_name":     user.Name,
				"user_hash":     user.Hash,
				"user_role":     user.RoleName,
				"provider_name": awc.prvname.String(),
				"provider_type": awc.prvtype.String(),
			})),
		),
	)
	assistantSpan := observation.Span(langfuse.WithSpanName("prepare assistant worker"))
	ctx, _ = assistantSpan.Observation(ctx)

	aslw, err := awc.aslc.NewFlowAssistantLog(ctx, awc.flowID, assistant.ID, pub)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to create flow assistant log worker", err)
	}

	prompter, err := newUserPrompter(ctx, awc.db, awc.userID)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to build user prompter", err)
	}
	executor, err := tools.NewFlowToolsExecutor(awc.db, awc.cfg, awc.docker, awc.functions, awc.userID, awc.flowID)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to create flow tools executor", err)
	}
	assistantProvider, err := awc.provs.NewAssistantProvider(ctx, awc.prvname, prompter, executor,
		assistant.ID, awc.flowID, awc.userID, container.Image, awc.input, aslw.StreamFlowAssistantMsg)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to get assistant provider", err)
	}

	msgChainID, err := assistantProvider.PrepareAgentChain(ctx)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to prepare assistant chain", err)
	}

	functionsBlob, err := json.Marshal(awc.functions)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to marshal functions", err)
	}

	logger = logger.WithField("msg_chain_id", msgChainID)
	logger.Info("assistant provider prepared")

	updated, err := awc.db.UpdateAssistant(ctx, database.UpdateAssistantParams{
		Title:              assistantProvider.Title(),
		Model:              assistantProvider.Model(pconfig.OptionsTypePrimaryAgent),
		Language:           assistantProvider.Language(),
		ToolCallIDTemplate: assistantProvider.ToolCallIDTemplate(),
		Functions:          functionsBlob,
		TraceID:            database.StringToNullString(observation.TraceID()),
		MsgchainID:         database.Int64ToNullInt64(&msgChainID),
		ID:                 assistant.ID,
	})
	if err != nil {
		logger.WithError(err).Error("failed to update assistant in DB")
		return nil, fmt.Errorf("failed to update assistant in DB: %w", err)
	}

	assistant = updated

	workers, err := getFlowProviderWorkers(ctx, awc.flowID, &awc.flowProviderControllers)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to get flow provider workers", err)
	}

	assistantProvider.SetAgentLogProvider(workers.alw)
	assistantProvider.SetMsgLogProvider(aslw)
	assistantProvider.SetFlowWorker(awc.fw)

	executor.SetImage(container.Image)
	executor.SetEmbedder(assistantProvider.Embedder())
	executor.SetScreenshotProvider(workers.sw)
	executor.SetAgentLogProvider(workers.alw)
	executor.SetMsgLogProvider(aslw)
	executor.SetSearchLogProvider(workers.slw)
	executor.SetTermLogProvider(workers.tlw)
	executor.SetVectorStoreLogProvider(workers.vslw)
	executor.SetToolCallLogProvider(workers.tclw)
	executor.SetKnowledgeProvider(pub)
	executor.SetGraphitiClient(awc.provs.GraphitiClient())

	ctx, cancel := context.WithCancel(context.Background())
	ctx, _ = obs.Observer.NewObservation(ctx, langfuse.WithObservationTraceID(observation.TraceID()))
	aw := &assistantWorker{
		id:      assistant.ID,
		flowID:  awc.flowID,
		userID:  awc.userID,
		chainID: msgChainID,
		aslw:    aslw,
		ap:      assistantProvider,
		db:      awc.db,
		wg:      &sync.WaitGroup{},
		pub:     pub,
		fw:      awc.fw,
		ctx:     ctx,
		cancel:  cancel,
		runMX:   &sync.Mutex{},
		runST:   func() {},
		runWG:   &sync.WaitGroup{},
		input:   make(chan assistantInput),
		logger: logrus.WithFields(logrus.Fields{
			"msg_chain_id": msgChainID,
			"assistant_id": assistant.ID,
			"flow_id":      awc.flowID,
			"user_id":      awc.userID,
			"trace_id":     observation.TraceID(),
			"component":    "assistant",
		}),
	}

	aw.wg.Add(1)
	go aw.worker()

	if err := awc.fw.AddAssistant(ctx, aw); err != nil {
		aw.cancel()
		aw.wg.Wait()

		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to attach assistant worker", err)
	}

	pub.AssistantUpdated(ctx, assistant)

	if err := aw.PutInput(ctx, awc.input, awc.useAgents, awc.resources); err != nil {
		if aw.ctx.Err() != nil || errors.Is(err, context.Canceled) {
			assistantSpan.End(langfuse.WithSpanStatus("assistant worker stopped or finished by its owner"))
			return aw, nil
		}

		aw.cancel()
		aw.wg.Wait()

		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to run assistant worker", err)
	}

	assistantSpan.End(langfuse.WithSpanStatus("assistant worker started"))

	return aw, nil
}

func LoadAssistantWorker(
	ctx context.Context, assistant database.Assistant, awc assistantWorkerCtx,
) (AssistantWorker, error) {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.LoadAssistantWorker")
	defer span.End()

	switch assistant.Status {
	case database.AssistantStatusRunning, database.AssistantStatusWaiting:
	default:
		return nil, fmt.Errorf("assistant %d has status %s: loading aborted: %w", assistant.ID, assistant.Status, ErrNothingToLoad)
	}

	logger := logrus.WithContext(ctx).WithFields(logrus.Fields{
		"assistant_id":  assistant.ID,
		"flow_id":       awc.flowID,
		"user_id":       awc.userID,
		"msg_chain_id":  assistant.MsgchainID,
		"provider_name": assistant.ModelProviderName,
		"provider_type": assistant.ModelProviderType,
	})

	user, err := awc.db.GetUser(ctx, awc.userID)
	if err != nil {
		logger.WithError(err).Error("failed to get user")
		return nil, fmt.Errorf("failed to get user %d: %w", awc.userID, err)
	}

	container, err := awc.db.GetFlowPrimaryContainer(ctx, awc.flowID)
	if err != nil {
		logger.WithError(err).Error("failed to get flow primary container")
		return nil, fmt.Errorf("failed to get flow primary container: %w", err)
	}

	ctx, observation := obs.Observer.NewObservation(ctx,
		langfuse.WithObservationTraceContext(
			langfuse.WithTraceName(fmt.Sprintf("%s%d flow %d assistant worker", awc.cfg.TenantLabel(), awc.flowID, assistant.ID)),
			langfuse.WithTraceUserID(tenantUserID(awc.cfg, user.Mail)),
			langfuse.WithTraceTags(tenantTags(awc.cfg, "controller", "assistant")),
			langfuse.WithTraceSessionID(awc.cfg.ScopedName(fmt.Sprintf("assistant-%d-flow-%d", assistant.ID, awc.flowID))),
			langfuse.WithTraceMetadata(tenantMeta(awc.cfg, langfuse.Metadata{
				"assistant_id":  assistant.ID,
				"flow_id":       awc.flowID,
				"user_id":       awc.userID,
				"user_email":    user.Mail,
				"user_name":     user.Name,
				"user_hash":     user.Hash,
				"user_role":     user.RoleName,
				"provider_name": assistant.ModelProviderName,
				"provider_type": assistant.ModelProviderType,
			})),
		),
	)
	assistantSpan := observation.Span(langfuse.WithSpanName("prepare assistant worker"))
	ctx, _ = assistantSpan.Observation(ctx)

	functions := &tools.Functions{}
	if err := json.Unmarshal(assistant.Functions, functions); err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to unmarshal functions", err)
	}

	pub := awc.subs.NewFlowPublisher(awc.userID, awc.flowID)
	aslw, err := awc.aslc.NewFlowAssistantLog(ctx, awc.flowID, assistant.ID, pub)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to create flow assistant log worker", err)
	}

	prompter := awc.prompter
	if prompter == nil {
		prompter, err = newUserPrompter(ctx, awc.db, awc.userID)
		if err != nil {
			return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to build user prompter", err)
		}
	}
	executor, err := tools.NewFlowToolsExecutor(awc.db, awc.cfg, awc.docker, functions, awc.userID, awc.flowID)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to create flow tools executor", err)
	}
	assistantProvider, err := awc.provs.LoadAssistantProvider(ctx, provider.ProviderName(assistant.ModelProviderName),
		prompter, executor, assistant.ID, awc.flowID, awc.userID, container.Image, assistant.Language, assistant.Title,
		assistant.ToolCallIDTemplate, aslw.StreamFlowAssistantMsg)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to get assistant provider", err)
	}

	workers, err := getFlowProviderWorkers(ctx, awc.flowID, &awc.flowProviderControllers)
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to get flow provider workers", err)
	}

	assistantProvider.SetAgentLogProvider(workers.alw)
	assistantProvider.SetMsgLogProvider(aslw)
	assistantProvider.SetFlowWorker(awc.fw)

	executor.SetImage(container.Image)
	executor.SetEmbedder(assistantProvider.Embedder())
	executor.SetScreenshotProvider(workers.sw)
	executor.SetAgentLogProvider(workers.alw)
	executor.SetMsgLogProvider(aslw)
	executor.SetSearchLogProvider(workers.slw)
	executor.SetTermLogProvider(workers.tlw)
	executor.SetVectorStoreLogProvider(workers.vslw)
	executor.SetToolCallLogProvider(workers.tclw)
	executor.SetKnowledgeProvider(pub)

	var msgChainID int64
	pmsgChainID := database.NullInt64ToInt64(assistant.MsgchainID)
	if pmsgChainID != nil {
		msgChainID = *pmsgChainID
		assistantProvider.SetMsgChainID(msgChainID)
	} else {
		return nil, fmt.Errorf("assistant %d has no msgchain id", assistant.ID)
	}

	ctx, cancel := context.WithCancel(context.Background())
	ctx, _ = obs.Observer.NewObservation(ctx, langfuse.WithObservationTraceID(observation.TraceID()))
	aw := &assistantWorker{
		id:      assistant.ID,
		flowID:  awc.flowID,
		userID:  awc.userID,
		chainID: msgChainID,
		aslw:    aslw,
		ap:      assistantProvider,
		db:      awc.db,
		wg:      &sync.WaitGroup{},
		pub:     pub,
		fw:      awc.fw,
		ctx:     ctx,
		cancel:  cancel,
		runMX:   &sync.Mutex{},
		runST:   func() {},
		runWG:   &sync.WaitGroup{},
		input:   make(chan assistantInput),
		logger: logrus.WithFields(logrus.Fields{
			"msg_chain_id": msgChainID,
			"assistant_id": assistant.ID,
			"flow_id":      awc.flowID,
			"user_id":      awc.userID,
			"trace_id":     observation.TraceID(),
			"component":    "assistant",
		}),
	}

	isInterrupted := assistant.Status == database.AssistantStatusRunning

	assistant, err = awc.db.UpdateAssistantStatus(ctx, database.UpdateAssistantStatusParams{
		Status: database.AssistantStatusWaiting,
		ID:     assistant.ID,
	})
	if err != nil {
		return nil, wrapErrorEndSpan(ctx, assistantSpan, "failed to update assistant status", err)
	}

	pub.AssistantUpdated(ctx, assistant)

	if isInterrupted {
		if _, err := aslw.PutFlowAssistantMsg(ctx, database.MsglogTypeReport, "", interruptedRunNote); err != nil {
			logger.WithError(err).Warn("failed to note the interrupted assistant run")
			aw.noteOwed.Store(true)
		}
	}

	aw.wg.Add(1)
	go aw.worker()

	assistantSpan.End(langfuse.WithSpanStatus("assistant worker started"))

	return aw, nil
}

func (aw *assistantWorker) noteInterruptedRun(ctx context.Context) {
	if !aw.noteOwed.Load() {
		return
	}

	if _, err := aw.aslw.PutFlowAssistantMsg(ctx, database.MsglogTypeReport, "", interruptedRunNote); err != nil {
		obs.LogErrorOrCancel(aw.logger, err, "failed to note the interrupted assistant run")

		return
	}

	aw.noteOwed.Store(false)
}

func (aw *assistantWorker) worker() {
	defer aw.wg.Done()

	perform := func(ctx context.Context, input string, useAgents bool) error {
		aw.runWG.Add(1)
		defer aw.runWG.Done()

		aw.noteInterruptedRun(ctx)

		_, err := aw.db.UpdateAssistantUseAgents(ctx, database.UpdateAssistantUseAgentsParams{
			UseAgents: useAgents,
			ID:        aw.id,
		})
		if err != nil {
			return fmt.Errorf("failed to update assistant use agents: %w", err)
		}

		if err := aw.SetStatus(ctx, database.AssistantStatusRunning); err != nil {
			obs.LogErrorOrCancel(aw.logger, err, "failed to set assistant status to waiting")
		}

		defer func() {
			if err := aw.SetStatus(ctx, database.AssistantStatusWaiting); err != nil {
				obs.LogErrorOrCancel(aw.logger, err, "failed to set assistant status to waiting")
			}
		}()

		_, err = aw.aslw.PutFlowAssistantMsg(ctx, database.MsglogTypeInput, "", input)
		if err != nil {
			return fmt.Errorf("failed to put input to flow assistant log: %w", err)
		}

		aw.runMX.Lock()
		ctx, aw.runST = context.WithCancel(aw.ctx)
		aw.runMX.Unlock()

		if err := aw.ap.PutInputToAgentChain(ctx, input); err != nil {
			return fmt.Errorf("failed to put input to agent chain: %w", err)
		}

		if err := aw.ap.PerformAgentChain(ctx); err != nil {
			if errors.Is(err, context.Canceled) {
				ctx = context.Background()
			}
			errChainConsistency := aw.ap.EnsureChainConsistency(ctx)
			if errChainConsistency != nil {
				err = errors.Join(err, errChainConsistency)
			}
			return fmt.Errorf("failed to perform agent chain: %w", err)
		}

		return nil
	}

	for {
		select {
		case <-aw.ctx.Done():
			return
		case ain := <-aw.input:
			err := perform(aw.ctx, ain.input, ain.useAgents)
			if err != nil {
				obs.LogErrorOrCancel(aw.logger, err, "failed to perform assistant chain")
			}
			if !ain.reply.deliver(err) {
				aw.reportFailure(err)
			}
		}
	}
}

func (aw *assistantWorker) reportFailure(err error) {
	if err == nil || errors.Is(err, context.Canceled) {
		return
	}

	msg := fmt.Sprintf("The assistant could not answer: %s", err)
	if _, putErr := aw.aslw.PutFlowAssistantMsg(aw.ctx, database.MsglogTypeReport, "", msg); putErr != nil {
		aw.logger.WithError(putErr).Warn("failed to report why the assistant did not answer")
	}
}

func (aw *assistantWorker) GetAssistantID() int64 {
	return aw.id
}

func (aw *assistantWorker) GetUserID() int64 {
	return aw.userID
}

func (aw *assistantWorker) GetFlowID() int64 {
	return aw.flowID
}

func (aw *assistantWorker) GetTitle() string {
	return aw.ap.Title()
}

func (aw *assistantWorker) GetStatus(ctx context.Context) (database.AssistantStatus, error) {
	assistant, err := aw.db.GetAssistant(ctx, aw.id)
	if err != nil {
		return database.AssistantStatusFailed, err
	}

	return assistant.Status, nil
}

func (aw *assistantWorker) SetStatus(ctx context.Context, status database.AssistantStatus) error {
	assistant, err := aw.db.UpdateAssistantStatus(ctx, database.UpdateAssistantStatusParams{
		Status: status,
		ID:     aw.id,
	})
	if err != nil {
		return fmt.Errorf("failed to update assistant %d flow %d status: %w", aw.id, aw.flowID, err)
	}

	aw.pub.AssistantUpdated(ctx, assistant)

	return nil
}

func (aw *assistantWorker) PutInput(ctx context.Context, input string, useAgents bool, resources []database.UserResource) error {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.assistantWorker.PutInput")
	defer span.End()

	if aw.fw != nil {
		if err := aw.fw.PutResources(ctx, resources); err != nil {
			aw.logger.WithError(err).Warn("failed to copy resources before assistant input")
		}
	}

	ain := assistantInput{input: input, useAgents: useAgents, reply: newInputReply()}
	select {
	case <-aw.ctx.Done():
		return fmt.Errorf("assistant %d flow %d stopped: %w", aw.id, aw.flowID, aw.ctx.Err())
	case <-ctx.Done():
		return fmt.Errorf("assistant %d flow %d input processing timeout: %w", aw.id, aw.flowID, ctx.Err())
	case aw.input <- ain:
		timer := time.NewTimer(assistantInputTimeout)
		defer timer.Stop()

		var stopErr error
		select {
		case err := <-ain.reply.done:
			return err
		case <-timer.C:
		case <-aw.ctx.Done():
			stopErr = fmt.Errorf("assistant %d flow %d stopped: %w", aw.id, aw.flowID, aw.ctx.Err())
		case <-ctx.Done():
			stopErr = fmt.Errorf("assistant %d flow %d input processing timeout: %w", aw.id, aw.flowID, ctx.Err())
		}

		if isDelivered, err := ain.reply.abandon(); isDelivered {
			return err
		}
		return stopErr
	}
}

func (aw *assistantWorker) Finish(ctx context.Context) error {
	ctx, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.assistantWorker.Finish")
	defer span.End()

	if err := aw.ctx.Err(); err != nil {
		if errors.Is(err, context.Canceled) {
			return nil
		}
		return fmt.Errorf("assistant %d flow %d stop failed: %w", aw.id, aw.flowID, err)
	}

	aw.cancel()
	aw.wg.Wait()

	if err := aw.SetStatus(ctx, database.AssistantStatusFinished); err != nil {
		aw.logger.WithError(err).Error("failed to set assistant status to finished")
	}

	return nil
}

func (aw *assistantWorker) Stop(ctx context.Context) error {
	_, span := obs.Observer.NewSpan(ctx, obs.SpanKindInternal, "controller.assistantWorker.Stop")
	defer span.End()

	aw.runST()
	done := make(chan struct{})
	timer := time.NewTimer(stopAssistantTimeout)
	defer timer.Stop()

	go func() {
		aw.runWG.Wait()
		close(done)
	}()

	select {
	case <-timer.C:
		return fmt.Errorf("assistant stop timeout")
	case <-done:
		return nil
	}
}
