transcript_processor.py 4.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151
  1. """
  2. Transcript Processor — deduplicates, stitches overlapping ASR windows,
  3. detects utterance boundaries, and formats streaming output for the client.
  4. """
  5. import time
  6. from dataclasses import dataclass, field
  7. from typing import Optional, List, Callable, Awaitable
  8. @dataclass
  9. class _Utterance:
  10. accumulated: str = ""
  11. last_raw: str = ""
  12. same_count: int = 0
  13. silence_since: float = 0.0
  14. class TranscriptProcessor:
  15. """
  16. Processes streaming ASR transcripts into clean, segmented utterances.
  17. Each call to feed() returns a list of zero or more message dicts:
  18. {"type": "partial", "text": "...", "seq": N}
  19. {"type": "utterance", "text": "...", "seq": N}
  20. """
  21. def __init__(
  22. self,
  23. silence_repeats: int = 3,
  24. silence_seconds: float = 3.0,
  25. post_hook: Optional[Callable[[str], Awaitable[str]]] = None,
  26. ):
  27. self.silence_repeats = silence_repeats
  28. self.silence_seconds = silence_seconds
  29. self._post_hook = post_hook
  30. self._current: Optional[_Utterance] = None
  31. self._global_seq = 0
  32. # ------------------------------------------------------------------
  33. # Public API
  34. # ------------------------------------------------------------------
  35. def feed(self, raw_text: str) -> List[dict]:
  36. """Feed a raw ASR transcript. Returns 0..N messages to send."""
  37. raw_text = raw_text.strip()
  38. if not raw_text:
  39. return []
  40. messages: List[dict] = []
  41. # ---- ensure we have an utterance context ----
  42. if self._current is None:
  43. self._current = _Utterance()
  44. # ---- 1. stitch ----
  45. if self._current.accumulated:
  46. merged, flush_msg = self._stitch(
  47. self._current.accumulated, raw_text
  48. )
  49. if flush_msg is not None:
  50. messages.append(flush_msg)
  51. # _stitch may have flushed (self._current = None). Restore.
  52. if self._current is None:
  53. self._current = _Utterance()
  54. else:
  55. merged = raw_text
  56. self._current.accumulated = merged
  57. # ---- 2. boundary detection ----
  58. if raw_text == self._current.last_raw:
  59. self._current.same_count += 1
  60. if self._current.same_count >= self.silence_repeats:
  61. if self._current.silence_since == 0.0:
  62. self._current.silence_since = time.monotonic()
  63. else:
  64. self._current.same_count = 0
  65. self._current.silence_since = 0.0
  66. self._current.last_raw = raw_text
  67. self._global_seq += 1
  68. messages.append({
  69. "type": "partial",
  70. "text": merged,
  71. "seq": self._global_seq,
  72. })
  73. # ---- 3. silence timer ----
  74. if (
  75. self._current.silence_since > 0
  76. and time.monotonic() - self._current.silence_since
  77. >= self.silence_seconds
  78. ):
  79. msg = self._flush()
  80. if msg:
  81. messages.append(msg)
  82. return messages
  83. def finalize(self) -> List[dict]:
  84. """Force-finalize (pause/end). Returns 0..1 messages."""
  85. return self._finish()
  86. # ------------------------------------------------------------------
  87. # Internals
  88. # ------------------------------------------------------------------
  89. def _stitch(
  90. self, previous: str, next_text: str
  91. ) -> tuple[str, Optional[dict]]:
  92. """
  93. Merge next_text into previous using longest suffix–prefix overlap.
  94. Returns (merged_text, pending_utterance_message_or_None).
  95. The pending message contains the OLD utterance if there is no
  96. overlap with previous — i.e. the ASR has started a new utterance.
  97. """
  98. if not next_text:
  99. return previous, None
  100. if next_text.startswith(previous):
  101. return next_text, None
  102. # Longest suffix–prefix overlap
  103. for n in range(len(next_text), 0, -1):
  104. needle = next_text[:n]
  105. if previous.endswith(needle):
  106. return previous + next_text[n:], None
  107. # No overlap at all — previous utterance is finished.
  108. flush = self._flush()
  109. return next_text, flush
  110. def _flush(self) -> Optional[dict]:
  111. """Flush current utterance. Returns utterance message or None."""
  112. if self._current is None:
  113. return None
  114. text = self._current.accumulated.strip()
  115. self._current = None
  116. if not text:
  117. return None
  118. self._global_seq += 1
  119. return {"type": "utterance", "text": text, "seq": self._global_seq}
  120. def _finish(self) -> List[dict]:
  121. """Finalize and clean up. Returns 0..1 messages."""
  122. messages: List[dict] = []
  123. msg = self._flush()
  124. if msg:
  125. messages.append(msg)
  126. self._current = None
  127. return messages