Skip to content
Merged
156 changes: 156 additions & 0 deletions backend/src/timeflow/gateway/websocket/agent_ports.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,156 @@
"""What the transport needs from a dialogue agent, stated on the transport's own terms."""

from collections.abc import AsyncIterator
from typing import Any, Protocol


class StreamIdentity(Protocol):
"""Identifiers of the audio stream a result belongs to."""

@property
def session_id(self) -> str:
"""Session the stream belongs to."""
...

@property
def stream_id(self) -> str:
"""The stream itself."""
...

@property
def conversation_id(self) -> str:
"""Conversation the stream continues."""
...

@property
def request_id(self) -> str | None:
"""Request that opened the stream, when the client supplied one."""
...


class TranscriptResult(Protocol):
"""What the user was heard to say."""

@property
def text(self) -> str:
"""The transcribed words."""
...

@property
def language(self) -> str:
"""Language the words were recognized as."""
...

@property
def duration_ms(self) -> int:
"""How long the audio ran."""
...


class ReplyTextProgress(Protocol):
"""How much of a reply's wording is known so far."""

@property
def reply_id(self) -> str:
"""Identifier tying one reply's run of updates together."""
...

@property
def speech_text(self) -> str:
"""Everything said so far, not only the newest part."""
...

@property
def done(self) -> bool:
"""Whether this is the last update for this reply."""
...


class DialogueQuestionInfo(Protocol):
"""A question the agent needs answered before it can act."""

@property
def question_id(self) -> str:
"""Identifier of this question."""
...

@property
def question_kind(self) -> str:
"""Why it is being asked."""
...

@property
def speech_text(self) -> str:
"""The question as it is spoken."""
...

@property
def required_response(self) -> str | None:
"""Field the answer should supply, when the question names one."""
...

@property
def candidates(self) -> tuple[dict[str, Any], ...]:
"""Choices the user is being asked to pick between, when there are any."""
...


class CommandOutcome(Protocol):
"""A command that was carried out, ready to be sent to the client."""

@property
def message_id(self) -> str:
"""Identifier the client echoes back in message.ack."""
...

@property
def operation(self) -> str:
"""Command that was carried out."""
...

@property
def status(self) -> str:
"""Outcome of that command."""
...

@property
def schedule(self) -> dict[str, Any]:
"""Persisted schedule snapshot the command produced."""
...


class AudioReplyInfo(Protocol):
"""Format and purpose of a spoken reply, announced before its audio."""

@property
def audio_id(self) -> str:
"""Identifier distinguishing this reply from the next."""
...

@property
def audio_format(self) -> str:
"""Encoding of the audio frames that follow."""
...

@property
def sample_rate_hz(self) -> int:
"""Sample rate the frames were produced at."""
...

@property
def purpose(self) -> str:
"""Why the reply is being spoken."""
...

@property
def speech_text(self) -> str:
"""The words the audio says."""
...


class Agent(Protocol):
"""Take one audio stream and act on what it contains."""

async def handle_audio(self, chunks: AsyncIterator[bytes], stream: StreamIdentity) -> None:
"""Consume the audio; returning only confirms receipt, not a result."""
...
39 changes: 39 additions & 0 deletions backend/src/timeflow/gateway/websocket/handlers/agent_audio.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
"""Audio sink that forwards an inbound stream to the dialogue agent unchanged."""

from collections.abc import AsyncIterator
from dataclasses import dataclass

from timeflow.gateway.websocket.agent_ports import Agent
from timeflow.gateway.websocket.ports import StreamContext


@dataclass(frozen=True, slots=True)
class _StreamIdentity:
"""Identifiers lifted out of a stream context for the agent."""

session_id: str
stream_id: str
conversation_id: str
request_id: str | None


class AgentAudioSink:
"""Hand each inbound stream to the agent as raw audio."""

def __init__(self, agent: Agent) -> None:
"""Store the agent that receives the audio."""
self._agent = agent

async def consume(self, chunks: AsyncIterator[bytes], stream: StreamContext) -> None:
"""Forward the stream unchanged, then return once the agent has taken it."""
await self._agent.handle_audio(chunks, _identity_of(stream))


def _identity_of(stream: StreamContext) -> _StreamIdentity:
"""Lift the identifiers the agent needs out of the transport's context."""
return _StreamIdentity(
session_id=stream.session.session_id,
stream_id=stream.stream_id,
conversation_id=stream.conversation_id,
request_id=stream.request_id,
)
142 changes: 142 additions & 0 deletions backend/src/timeflow/gateway/websocket/handlers/agent_result.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,142 @@
"""Result sink translating each of the agent's outlets into the protocol message for it."""

import logging
from collections.abc import AsyncIterator

from timeflow.gateway.websocket.agent_ports import (
AudioReplyInfo,
CommandOutcome,
DialogueQuestionInfo,
ReplyTextProgress,
StreamIdentity,
TranscriptResult,
)
from timeflow.gateway.websocket.connection_manager import ConnectionManager
from timeflow.gateway.websocket.messages.agent import (
VoiceAsrCompleted,
VoiceAsrCompletedPayload,
VoiceCommandResult,
VoiceCommandResultPayload,
)
from timeflow.gateway.websocket.messages.dialogue import (
QUESTION_KINDS,
QuestionKind,
VoiceDialogueQuestion,
VoiceDialogueQuestionPayload,
VoiceDialogueReply,
VoiceDialogueReplyPayload,
)
from timeflow.gateway.websocket.messages.tts import (
VoiceTtsEnd,
VoiceTtsStart,
VoiceTtsStartPayload,
)

logger = logging.getLogger(__name__)


def _question_kind(value: str) -> QuestionKind:
"""Narrow a producer's string to the protocol's four kinds, refusing anything else."""
if value not in QUESTION_KINDS:
raise ValueError(f"question_kind must be one of {QUESTION_KINDS}, got {value!r}")
return value


class WebSocketResultSink:
"""Translate results into wire messages and push them to the session."""

def __init__(self, connections: ConnectionManager) -> None:
"""Store the registry used to reach the session."""
self._connections = connections

async def deliver_transcript(
self, transcript: TranscriptResult, stream: StreamIdentity
) -> None:
"""Push what the user was heard to say."""
message = VoiceAsrCompleted(
request_id=stream.request_id,
conversation_id=stream.conversation_id,
payload=VoiceAsrCompletedPayload(
transcript=transcript.text,
language=transcript.language,
duration_ms=transcript.duration_ms,
),
)
await self._send(stream.session_id, message.type, message.model_dump())

async def deliver_reply_text(self, reply: ReplyTextProgress, stream: StreamIdentity) -> None:
"""Push how much of the reply's wording is known so far."""
message = VoiceDialogueReply(
request_id=stream.request_id,
conversation_id=stream.conversation_id,
payload=VoiceDialogueReplyPayload(
reply_id=reply.reply_id,
speech_text=reply.speech_text,
done=reply.done,
),
)
await self._send(stream.session_id, message.type, message.model_dump())

async def deliver_result(self, result: CommandOutcome, stream: StreamIdentity) -> None:
"""Push the outcome of a command the agent carried out."""
message = VoiceCommandResult(
message_id=result.message_id,
request_id=stream.request_id,
conversation_id=stream.conversation_id,
payload=VoiceCommandResultPayload(
operation=result.operation,
status=result.status,
schedule=result.schedule,
),
)
await self._send(stream.session_id, message.type, message.model_dump())

async def deliver_question(
self, question: DialogueQuestionInfo, stream: StreamIdentity
) -> None:
"""Push a question the user has to answer before the turn can go further."""
message = VoiceDialogueQuestion(
request_id=stream.request_id,
conversation_id=stream.conversation_id,
payload=VoiceDialogueQuestionPayload(
question_id=question.question_id,
question_kind=_question_kind(question.question_kind),
speech_text=question.speech_text,
required_response=question.required_response,
candidates=list(question.candidates),
),
)
await self._send(stream.session_id, message.type, message.model_dump())

async def deliver_audio(
self, reply: AudioReplyInfo, chunks: AsyncIterator[bytes], stream: StreamIdentity
) -> None:
"""Speak a reply, putting each chunk on the wire as it is produced."""
start = VoiceTtsStart(
conversation_id=stream.conversation_id,
audio_id=reply.audio_id,
payload=VoiceTtsStartPayload(
format=reply.audio_format,
sample_rate_hz=reply.sample_rate_hz,
purpose=reply.purpose,
speech_text=reply.speech_text,
),
)
end = VoiceTtsEnd(conversation_id=stream.conversation_id, audio_id=reply.audio_id)

delivered = await self._connections.stream_audio(
stream.session_id, start.model_dump(), chunks, end.model_dump()
)
if not delivered:
logger.info(
"stopped speaking to a session that had gone",
extra={"session_id": stream.session_id, "audio_id": reply.audio_id},
)

async def _send(self, session_id: str, message_type: str, envelope: dict[str, object]) -> None:
"""Send one message, logging rather than raising when the session has gone."""
if not await self._connections.send(session_id, envelope):
logger.info(
"dropped a result for a session that had gone",
extra={"session_id": session_id, "message_type": message_type},
)
34 changes: 34 additions & 0 deletions backend/src/timeflow/gateway/websocket/handlers/message_ack.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,34 @@
"""The message.ack confirmation, recorded and never answered."""

import logging
from typing import Any

from pydantic import ValidationError

from timeflow.gateway.websocket.envelope import ERROR_MALFORMED_MESSAGE, build_error_envelope
from timeflow.gateway.websocket.messages.agent import MessageAck
from timeflow.gateway.websocket.ports import SessionContext

logger = logging.getLogger(__name__)


async def handle_message_ack(
raw_message: dict[str, Any], session: SessionContext
) -> dict[str, Any] | None:
"""Record an applied command result and reply with nothing."""
try:
ack = MessageAck.model_validate(raw_message)
except ValidationError:
return build_error_envelope(
"message.ack", None, ERROR_MALFORMED_MESSAGE, "message.ack payload is invalid"
)

logger.info(
"client acknowledged a command result",
extra={
"session_id": session.session_id,
"message_id": ack.message_id,
"status": ack.status,
},
)
return None
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@ async def handle_start(
conversation_id=payload.conversation_id or self._conversation_id_factory(),
session=session,
audio_config=audio_config,
request_id=request_id,
)
stream = _ActiveStream(
context=context,
Expand Down
Loading
Loading