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
4 changes: 4 additions & 0 deletions invokeai/app/invocations/baseinvocation.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,10 @@ def invoke(self, context: InvocationContext) -> BaseInvocationOutput:
"""Invoke with provided context and return outputs."""
pass

def get_event_invocation(self) -> "BaseInvocation":
"""Returns the invocation representation included in execution events."""
return self

def invoke_internal(self, context: InvocationContext, services: "InvocationServices") -> BaseInvocationOutput:
"""
Internal invoke method, calls `invoke()` after some prep.
Expand Down
18 changes: 14 additions & 4 deletions invokeai/app/services/events/events_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ def dispatch(self, event: "EventBase") -> None:

def emit_invocation_started(self, queue_item: "SessionQueueItem", invocation: "BaseInvocation") -> None:
"""Emitted when an invocation is started"""
self.dispatch(InvocationStartedEvent.build(queue_item, invocation))
self.dispatch(InvocationStartedEvent.build(queue_item, invocation.get_event_invocation()))

def emit_invocation_progress(
self,
Expand All @@ -75,13 +75,15 @@ def emit_invocation_progress(
image: "ProgressImage | None" = None,
) -> None:
"""Emitted at periodically during an invocation"""
self.dispatch(InvocationProgressEvent.build(queue_item, invocation, message, percentage, image))
self.dispatch(
InvocationProgressEvent.build(queue_item, invocation.get_event_invocation(), message, percentage, image)
)

def emit_invocation_complete(
self, queue_item: "SessionQueueItem", invocation: "BaseInvocation", output: "BaseInvocationOutput"
) -> None:
"""Emitted when an invocation is complete"""
self.dispatch(InvocationCompleteEvent.build(queue_item, invocation, output))
self.dispatch(InvocationCompleteEvent.build(queue_item, invocation.get_event_invocation(), output))

def emit_invocation_error(
self,
Expand All @@ -92,7 +94,15 @@ def emit_invocation_error(
error_traceback: str,
) -> None:
"""Emitted when an invocation encounters an error"""
self.dispatch(InvocationErrorEvent.build(queue_item, invocation, error_type, error_message, error_traceback))
self.dispatch(
InvocationErrorEvent.build(
queue_item,
invocation.get_event_invocation(),
error_type,
error_message,
error_traceback,
)
)

# endregion

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@
WorkflowCallQueueLifecycle,
)
from invokeai.app.services.session_queue.session_queue_common import SessionQueueItem, SessionQueueItemNotFoundError
from invokeai.app.services.shared.graph import NodeInputError
from invokeai.app.services.shared.graph import CollectInvocation, IterateInvocation, NodeInputError
from invokeai.app.services.shared.invocation_context import InvocationContextData, build_invocation_context
from invokeai.app.util.profiler import Profiler

Expand Down Expand Up @@ -141,10 +141,19 @@ def run_node(self, invocation: BaseInvocation, queue_item: SessionQueueItem):

# Invoke the node
output = invocation.invoke_internal(context=context, services=self._services)
control_collection = None
if self._on_after_run_node_callbacks and isinstance(invocation, (IterateInvocation, CollectInvocation)):
control_collection = invocation.collection
# Save output and history
queue_item.session.complete(invocation.id, output)

self._on_after_run_node(invocation, queue_item, output)
if control_collection is not None:
invocation.collection = control_collection
try:
self._on_after_run_node(invocation, queue_item, output)
finally:
if control_collection is not None:
invocation.collection = []

except CanceledException:
# A CanceledException is raised during the denoising step callback if the cancel event is set. We don't need
Expand Down Expand Up @@ -215,7 +224,7 @@ def _on_after_run_session(self, queue_item: SessionQueueItem) -> None:
# The queue item may have been canceled or failed while the session was running. We should only complete it
# if it is not already canceled or failed.
if queue_item.status not in ["canceled", "failed"] and queue_item.session.is_complete():
queue_item = self._services.session_queue.complete_queue_item(queue_item.item_id)
queue_item = self._services.session_queue.complete_queue_item(queue_item.item_id, queue_item=queue_item)

# We'll get a GESStatsNotFoundError if we try to log stats for an untracked graph, but in the processor
# we don't care about that - suppress the error.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -96,7 +96,7 @@ def begin_workflow_call_boundary(
child_queue_item = None
enqueued_child_item_ids: list[int] = []
try:
self._session_runner._services.session_queue.set_queue_item_session(queue_item.item_id, queue_item.session)
self._session_runner._services.session_queue.save_queue_item_session(queue_item.item_id, queue_item.session)
for child_result in child_session_results:
child_queue_item = self._session_runner._services.session_queue.enqueue_workflow_call_child(
parent_queue_item=queue_item,
Expand All @@ -105,8 +105,8 @@ def begin_workflow_call_boundary(
)
enqueued_child_item_ids.append(child_queue_item.item_id)
queue_item.session.set_waiting_workflow_call_child_item_ids(enqueued_child_item_ids)
self._session_runner._services.session_queue.set_queue_item_session(queue_item.item_id, queue_item.session)
self._session_runner._services.session_queue.suspend_queue_item(queue_item.item_id)
self._session_runner._services.session_queue.save_queue_item_session(queue_item.item_id, queue_item.session)
self._session_runner._services.session_queue.suspend_queue_item(queue_item.item_id, queue_item=queue_item)
except Exception as e:
if enqueued_child_item_ids:
self._session_runner._services.session_queue.delete_queue_items_by_id(enqueued_child_item_ids)
Expand Down Expand Up @@ -218,7 +218,7 @@ def _resume_parent_from_completed_child(self, child_queue_item: SessionQueueItem
self._fail_parent_from_failed_child(parent_queue_item)
return
if not should_resume_parent:
self._session_runner._services.session_queue.set_queue_item_session(
self._session_runner._services.session_queue.save_queue_item_session(
parent_queue_item.item_id, parent_queue_item.session
)
return
Expand All @@ -228,17 +228,19 @@ def _resume_parent_from_completed_child(self, child_queue_item: SessionQueueItem
parent_output = WorkflowReturnOutput(values=aggregated_values)
parent_queue_item.session.complete(waiting_invocation.id, parent_output)
self._session_runner._on_after_run_node(waiting_invocation, parent_queue_item, parent_output)
parent_queue_item = self._session_runner._services.session_queue.set_queue_item_session(
self._session_runner._services.session_queue.save_queue_item_session(
parent_queue_item.item_id, parent_queue_item.session
)
if parent_queue_item.session.is_complete():
parent_queue_item = self._session_runner._services.session_queue.complete_queue_item(
parent_queue_item.item_id
parent_queue_item.item_id, queue_item=parent_queue_item
)
if getattr(parent_queue_item, "parent_item_id", None) is not None:
self._resume_parent_from_completed_child(parent_queue_item)
return
self._session_runner._services.session_queue.resume_queue_item(parent_queue_item.item_id)
self._session_runner._services.session_queue.resume_queue_item(
parent_queue_item.item_id, queue_item=parent_queue_item
)

def _fail_parent_from_failed_child(self, child_queue_item: SessionQueueItem) -> None:
parent_queue_item = self._get_parent_queue_item(child_queue_item)
Expand Down
15 changes: 10 additions & 5 deletions invokeai/app/services/session_queue/session_queue_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,8 +92,8 @@ def get_queue_status(
acting_user_id is independent of user_id and controls only current-item redaction:
when set, the returned status omits item_id/session_id/batch_id unless the
currently-running item belongs to acting_user_id. The redaction is decided from the
same get_current() snapshot used to embed those identifiers, so it cannot race against
a concurrent state change.
same database snapshot used to embed those identifiers, so it cannot race against a
concurrent state change.

is_admin disables current-item redaction entirely: admins may see the identifiers of
any user's current item. Redaction stays fail-closed - a caller that passes user_id
Expand All @@ -114,17 +114,17 @@ def get_batch_status(self, queue_id: str, batch_id: str, user_id: Optional[str]
pass

@abstractmethod
def complete_queue_item(self, item_id: int) -> SessionQueueItem:
def complete_queue_item(self, item_id: int, queue_item: Optional[SessionQueueItem] = None) -> SessionQueueItem:
"""Completes a session queue item"""
pass

@abstractmethod
def suspend_queue_item(self, item_id: int) -> SessionQueueItem:
def suspend_queue_item(self, item_id: int, queue_item: Optional[SessionQueueItem] = None) -> SessionQueueItem:
"""Suspends a session queue item while waiting on a child workflow execution."""
pass

@abstractmethod
def resume_queue_item(self, item_id: int) -> SessionQueueItem:
def resume_queue_item(self, item_id: int, queue_item: Optional[SessionQueueItem] = None) -> SessionQueueItem:
"""Resumes a suspended session queue item by returning it to pending state."""
pass

Expand Down Expand Up @@ -228,6 +228,11 @@ def set_queue_item_session(self, item_id: int, session: GraphExecutionState) ->
"""Sets the session for a session queue item. Use this to update the session state."""
pass

@abstractmethod
def save_queue_item_session(self, item_id: int, session: GraphExecutionState) -> None:
"""Persists a queue item's session without loading and returning the full queue item."""
pass

@abstractmethod
def enqueue_workflow_call_child(
self,
Expand Down
Loading
Loading