Skip to content
Merged
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
2 changes: 2 additions & 0 deletions backend/druks/chat/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
)
# The line between what a person typed and what they said in the same message.
VOICE_NOTE_MARKER = "[Voice note]"
# The person reads this, so it holds no error details.
DELIVERY_FAILED_MESSAGE = "The reply is unavailable. Send the message again to retry."
TRANSCRIPTION_FAILED_MESSAGE = (
"[Internal: Druks could not turn the person's voice note into text. Tell the person "
"in one short line to write it instead.]"
Expand Down
1 change: 1 addition & 0 deletions backend/druks/chat/enums.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ class MessageState(StrEnum):
REPLIED = "replied"
INTERRUPTED = "interrupted"
CANCELLED = "cancelled"
FAILED = "failed"


class PauseSignal(StrEnum):
Expand Down
13 changes: 13 additions & 0 deletions backend/druks/chat/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -327,6 +327,19 @@ async def end_turn(self, session: AsyncSession, state: MessageState) -> None:
.execution_options(synchronize_session="fetch")
)

async def end_pending_messages(
self, session: AsyncSession, state: MessageState
) -> list[Message]:
"""Give every pending message its final state, and return those messages."""
return list(
await session.scalars(
update(Message)
.where(Message.conversation_id == self.id, Message.state == MessageState.PENDING)
.values(state=state)
.returning(Message)
)
)

async def is_held(self, session: AsyncSession) -> bool:
"""Whether the chat's turns wait: a person answers it from the connection's
phone, or its connection was removed."""
Expand Down
41 changes: 41 additions & 0 deletions backend/druks/chat/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from druks.api.dependencies import SessionDep
from druks.apps.registry import channels

from .enums import MessageState
from .exceptions import ChannelHasNoThreadsError, ChatSandboxGone
from .models import Conversation, Message
from .schemas import ConversationDetailResponse, ConversationResponse, MessageResponse
Expand Down Expand Up @@ -111,6 +112,46 @@ async def create_message(
return message


@router.post(
"/conversations/{conversation_id}/messages/{message_id}/retry",
status_code=202,
response_model=MessageResponse,
response_model_by_alias=True,
)
async def retry_message(
conversation_id: str,
message_id: str,
session: SessionDep,
account: Account = Depends(current_session_account),
) -> Message:
conversation = await Conversation.get_for_account(session, conversation_id, account.id)
if not conversation:
raise HTTPException(404, "Conversation not found.")
original_message = await conversation.get_message(session, message_id)
if not original_message:
raise HTTPException(404, "Message not found.")
if original_message.state not in (
MessageState.FAILED,
MessageState.INTERRUPTED,
MessageState.CANCELLED,
):
raise HTTPException(
409,
f"The message is {original_message.state}. Send again only a failed, "
"interrupted, or cancelled message.",
)
message = await conversation.create_message(
session,
original_message.body,
file=original_message.file,
is_internal=original_message.is_internal,
)
await session.commit()
await publish(conversation.id, {"type": "messages"})
await DBOS.start_workflow_async(deliver, conversation.id)
return message


@router.post("/conversations/{conversation_id}/cancel", status_code=204)
async def cancel(
conversation_id: str,
Expand Down
202 changes: 120 additions & 82 deletions backend/druks/chat/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
BRIDGE_SETTLED_STATUSES,
CHAT_KEY_NAME,
CONVERSATION_HEADER,
DELIVERY_FAILED_MESSAGE,
FAILURE_MESSAGE,
INTERNAL_MESSAGES_PROMPT,
RESULT_MESSAGE,
Expand Down Expand Up @@ -196,98 +197,135 @@ async def deliver_turns() -> None:
conversation = await session.get(Conversation, conversation_id)
await deliver_pending(session, conversation)

try:
await DBOS.run_step_async(
StepOptions(
name="chat.deliver.turns",
retries_allowed=True,
max_attempts=5,
should_retry=lambda error: getattr(error, "is_retryable", True),
),
deliver_turns,
)
except ChatHarnessError as error:
await publish(conversation_id, {"type": "error", "detail": str(error)})
except Exception:
logger.exception("Chat delivery failed for conversation %s", conversation_id)
detail = "The reply is unavailable. Connect again to retry."
async def end_turns() -> str | None:
"""Give the pending messages the failed state; a delivered turn runs on. On a
live channel, also save a reply that says so, and return its id."""
async with step_session() as session:
conversation = await session.get(Conversation, conversation_id)
messages = await conversation.end_pending_messages(session, MessageState.FAILED)
person_messages = [message for message in messages if not message.is_internal]
if (
person_messages
and conversation.connection
and not await conversation.is_held(session)
):
newest_message = max(
person_messages, key=lambda message: (message.created_at, message.id)
)
reply = await conversation.create_message(
session,
DELIVERY_FAILED_MESSAGE,
role=MessageRole.ASSISTANT,
reply_to=newest_message,
)
return reply.id

async def send_reply(reply_id: str) -> None:
async with step_session() as session:
conversation = await session.get(Conversation, conversation_id)
reply = await session.get(Message, reply_id)
await channels.get(conversation.source).send_reply(session, conversation, reply)

async def record_failed(detail: str) -> None:
reply_id = await DBOS.run_step_async(StepOptions(name="chat.deliver.failed"), end_turns)
await reset_live_stream(conversation_id)
await publish(conversation_id, {"type": "error", "detail": detail})
if reply_id:
await DBOS.run_step_async(
StepOptions(name="chat.deliver.reply", retries_allowed=True, max_attempts=5),
send_reply,
reply_id,
)

# The lock covers the retries and the failure. Another delivery cannot take the
# messages between them.
async with lock(f"chat:{conversation_id}:delivery"):
try:
await DBOS.run_step_async(
StepOptions(
name="chat.deliver.turns",
retries_allowed=True,
max_attempts=5,
should_retry=lambda error: getattr(error, "is_retryable", True),
),
deliver_turns,
)
except ChatHarnessError as error:
await record_failed(f"{error} Send the message again to retry.")
except Exception:
logger.exception("Chat delivery failed for conversation %s", conversation_id)
await record_failed(DELIVERY_FAILED_MESSAGE)


async def deliver_pending(session: AsyncSession, conversation: Conversation) -> None:
await session.commit()
async with lock(f"chat:{conversation.id}:delivery"):
while message := await conversation.get_unanswered_message(session):
if message.state == MessageState.PENDING:
if await conversation.is_held(session):
return
await reset_live_stream(conversation.id)
config, prompt, tools = await get_agent(session, conversation)
if not config.harness_class.adapter_command:
adapters = ", ".join(
harness.name for harness in get_harnesses() if harness.adapter_command
)
raise ChatHarnessError(
f"Chat runs on {adapters}. The Bot's settings select "
f"{config.harness_class.name}. Set its harness to one of them."
)
# A bot serves outside people, so only the operator reaches the enabled MCP
# servers and holds their own sign-ins.
mcp_servers, secret_refs = (), []
if conversation.account.kind == AccountKind.OPERATOR:
# A server the operator has not connected must not stop the chat.
mcp_servers, secret_refs = await Workspace.get_all_mcp_servers(
session, None, conversation.account_id, skip_unauthenticated=True
)
for service in services.all():
if service.host and service.token_endpoint:
sign_ins = await VaultSecret.list_account_connections(
session, Audience.service(service.slug), conversation.account_id
)
# A sandbox has one variable per service, so it holds one sign-in.
secret_refs += [
SecretRef(
name=service.slug, secret_id=sign_in.id, host=service.host
)
for sign_in in sign_ins[:1]
]
host, identity = await get_sandbox(
while message := await conversation.get_unanswered_message(session):
if message.state == MessageState.PENDING:
if await conversation.is_held(session):
return
await reset_live_stream(conversation.id)
config, prompt, tools = await get_agent(session, conversation)
if not config.harness_class.adapter_command:
adapters = ", ".join(
harness.name for harness in get_harnesses() if harness.adapter_command
)
raise ChatHarnessError(
f"Chat runs on {adapters}. The Bot's settings select "
f"{config.harness_class.name}. Set its harness to one of them."
)
# A bot serves outside people, so only the operator reaches the enabled MCP
# servers and holds their own sign-ins.
mcp_servers, secret_refs = (), []
if conversation.account.kind == AccountKind.OPERATOR:
# A server the operator has not connected must not stop the chat.
mcp_servers, secret_refs = await Workspace.get_all_mcp_servers(
session, None, conversation.account_id, skip_unauthenticated=True
)
for service in services.all():
if service.host and service.token_endpoint:
sign_ins = await VaultSecret.list_account_connections(
session, Audience.service(service.slug), conversation.account_id
)
# A sandbox has one variable per service, so it holds one sign-in.
secret_refs += [
SecretRef(name=service.slug, secret_id=sign_in.id, host=service.host)
for sign_in in sign_ins[:1]
]
host, identity = await get_sandbox(
session,
conversation.account_id,
config=config,
allowed_tools=tools,
secret_refs=secret_refs,
)
try:
bridge = Bridge(host)
turn = await send_turn(
session,
conversation.account_id,
conversation,
message,
bridge=bridge,
identity=identity,
config=config,
allowed_tools=tools,
secret_refs=secret_refs,
prompt=prompt,
mcp_servers=mcp_servers,
)
try:
bridge = Bridge(host)
turn = await send_turn(
session,
conversation,
message,
bridge=bridge,
identity=identity,
config=config,
prompt=prompt,
mcp_servers=mcp_servers,
)
if turn:
await follow_turn(session, conversation, turn, bridge)
finally:
await host.aclose()
continue
try:
host = await get_running_sandbox(session, conversation.account_id)
except ChatSandboxGone:
await conversation.end_turn(session, MessageState.INTERRUPTED)
await session.commit()
await reset_live_stream(conversation.id)
continue
# The snapshot this asks for replaces an error a failed delivery left in the stream.
await publish(conversation.id, {"type": "messages"})
try:
await follow_turn(session, conversation, message, Bridge(host))
if turn:
await follow_turn(session, conversation, turn, bridge)
finally:
await host.aclose()
continue
try:
host = await get_running_sandbox(session, conversation.account_id)
except ChatSandboxGone:
await conversation.end_turn(session, MessageState.INTERRUPTED)
await session.commit()
await reset_live_stream(conversation.id)
continue
try:
await follow_turn(session, conversation, message, Bridge(host))
finally:
await host.aclose()


async def send_turn(
Expand Down
2 changes: 2 additions & 0 deletions backend/tests/test_auth_boundary.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
("GET", "/api/chat/conversations/{conversation_id}"),
("PATCH", "/api/chat/conversations/{conversation_id}"),
("POST", "/api/chat/conversations/{conversation_id}/messages"),
("POST", "/api/chat/conversations/{conversation_id}/messages/{message_id}/retry"),
("POST", "/api/chat/conversations/{conversation_id}/cancel"),
("GET", "/api/chat/services/waha/sessions"),
("POST", "/api/chat/services/waha/sessions"),
Expand Down Expand Up @@ -68,6 +69,7 @@
"/api/chat/conversations",
"/api/chat/conversations/{conversation_id}",
"/api/chat/conversations/{conversation_id}/messages",
"/api/chat/conversations/{conversation_id}/messages/{message_id}/retry",
"/api/chat/conversations/{conversation_id}/cancel",
"/api/chat/services/waha/sessions",
"/api/chat/services/waha/sessions/{session_id}/qr",
Expand Down
2 changes: 1 addition & 1 deletion backend/tests/test_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -354,7 +354,7 @@ async def request(self, method, **values):
if method == "status":
statuses += 1
if statuses > 1:
assert await get_client().xlen(service.events_key(conversation.id)) == 5
assert await get_client().xlen(service.events_key(conversation.id)) == len(events)
return {
"status": "running" if statuses == 1 else terminal_state,
"messageId": message.id,
Expand Down
Loading
Loading