Skip to content
Closed
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
108 changes: 108 additions & 0 deletions backend/src/timeflow/gateway/websocket/agent_ports.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
"""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 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,
)
95 changes: 95 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,95 @@
"""Result sink that pushes transcripts, command results and spoken replies to the client."""

import logging
from collections.abc import AsyncIterator

from timeflow.gateway.websocket.agent_ports import (
AudioReplyInfo,
CommandOutcome,
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.tts import (
VoiceTtsEnd,
VoiceTtsStart,
VoiceTtsStartPayload,
)

logger = logging.getLogger(__name__)


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_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_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
48 changes: 48 additions & 0 deletions backend/src/timeflow/gateway/websocket/messages/agent.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
"""Transcript and command result messages pushed after a stream ends."""

from typing import Any, Literal

from pydantic import BaseModel


class VoiceAsrCompletedPayload(BaseModel):
"""What the user was heard to say."""

transcript: str
language: str
duration_ms: int


class VoiceAsrCompleted(BaseModel):
"""Server message carrying the final transcript of a stream."""

type: Literal["voice.asr.completed"] = "voice.asr.completed"
request_id: str | None = None
conversation_id: str
payload: VoiceAsrCompletedPayload


class VoiceCommandResultPayload(BaseModel):
"""The command that was carried out and the schedule it produced."""

operation: str
status: str
schedule: dict[str, Any]


class VoiceCommandResult(BaseModel):
"""Server message carrying a committed command result."""

type: Literal["voice.command.result"] = "voice.command.result"
message_id: str
request_id: str | None = None
conversation_id: str
payload: VoiceCommandResultPayload


class MessageAck(BaseModel):
"""Client confirmation that a command result was applied locally."""

type: Literal["message.ack"]
message_id: str
status: str
1 change: 1 addition & 0 deletions backend/src/timeflow/gateway/websocket/ports.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ class StreamContext:
conversation_id: str
session: SessionContext
audio_config: AudioConfig
request_id: str | None = None


class TokenVerifier(Protocol):
Expand Down
Loading