BaseTextFilter: make functions async
This commit is contained in:
@@ -68,6 +68,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Changed
|
### Changed
|
||||||
|
|
||||||
|
- `BaseTextFilter` methods `filter()`, `update_settings()`,
|
||||||
|
`handle_interruption()` and `reset_interruption()` are now async.
|
||||||
|
|
||||||
- `BaseTextAggregator` methods `aggregate()`, `handle_interruption()` and
|
- `BaseTextAggregator` methods `aggregate()`, `handle_interruption()` and
|
||||||
`reset()` are now async.
|
`reset()` are now async.
|
||||||
|
|
||||||
|
|||||||
@@ -188,7 +188,7 @@ class HostResponseTextFilter(BaseTextFilter):
|
|||||||
# No settings to update for this filter
|
# No settings to update for this filter
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def filter(self, text: str) -> str:
|
async def filter(self, text: str) -> str:
|
||||||
# Remove case and whitespace for comparison
|
# Remove case and whitespace for comparison
|
||||||
clean_text = text.strip().upper()
|
clean_text = text.strip().upper()
|
||||||
|
|
||||||
@@ -198,10 +198,10 @@ class HostResponseTextFilter(BaseTextFilter):
|
|||||||
|
|
||||||
return text
|
return text
|
||||||
|
|
||||||
def handle_interruption(self):
|
async def handle_interruption(self):
|
||||||
self._interrupted = True
|
self._interrupted = True
|
||||||
|
|
||||||
def reset_interruption(self):
|
async def reset_interruption(self):
|
||||||
self._interrupted = False
|
self._interrupted = False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -178,7 +178,7 @@ class HostResponseTextFilter(BaseTextFilter):
|
|||||||
# No settings to update for this filter
|
# No settings to update for this filter
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def filter(self, text: str) -> str:
|
async def filter(self, text: str) -> str:
|
||||||
# Remove case and whitespace for comparison
|
# Remove case and whitespace for comparison
|
||||||
clean_text = text.strip().upper()
|
clean_text = text.strip().upper()
|
||||||
|
|
||||||
@@ -188,10 +188,10 @@ class HostResponseTextFilter(BaseTextFilter):
|
|||||||
|
|
||||||
return text
|
return text
|
||||||
|
|
||||||
def handle_interruption(self):
|
async def handle_interruption(self):
|
||||||
self._interrupted = True
|
self._interrupted = True
|
||||||
|
|
||||||
def reset_interruption(self):
|
async def reset_interruption(self):
|
||||||
self._interrupted = False
|
self._interrupted = False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -157,7 +157,7 @@ class TTSService(AIService):
|
|||||||
self.set_voice(value)
|
self.set_voice(value)
|
||||||
elif key == "text_filter":
|
elif key == "text_filter":
|
||||||
for filter in self._text_filters:
|
for filter in self._text_filters:
|
||||||
filter.update_settings(value)
|
await filter.update_settings(value)
|
||||||
else:
|
else:
|
||||||
logger.warning(f"Unknown setting for TTS service: {key}")
|
logger.warning(f"Unknown setting for TTS service: {key}")
|
||||||
|
|
||||||
@@ -236,7 +236,7 @@ class TTSService(AIService):
|
|||||||
self._processing_text = False
|
self._processing_text = False
|
||||||
await self._text_aggregator.handle_interruption()
|
await self._text_aggregator.handle_interruption()
|
||||||
for filter in self._text_filters:
|
for filter in self._text_filters:
|
||||||
filter.handle_interruption()
|
await filter.handle_interruption()
|
||||||
|
|
||||||
async def _maybe_pause_frame_processing(self):
|
async def _maybe_pause_frame_processing(self):
|
||||||
if self._processing_text and self._pause_frame_processing:
|
if self._processing_text and self._pause_frame_processing:
|
||||||
@@ -274,8 +274,8 @@ class TTSService(AIService):
|
|||||||
|
|
||||||
# Process all filter.
|
# Process all filter.
|
||||||
for filter in self._text_filters:
|
for filter in self._text_filters:
|
||||||
filter.reset_interruption()
|
await filter.reset_interruption()
|
||||||
text = filter.filter(text)
|
text = await filter.filter(text)
|
||||||
|
|
||||||
if text:
|
if text:
|
||||||
await self.process_generator(self.run_tts(text))
|
await self.process_generator(self.run_tts(text))
|
||||||
|
|||||||
@@ -10,17 +10,17 @@ from typing import Any, Mapping
|
|||||||
|
|
||||||
class BaseTextFilter(ABC):
|
class BaseTextFilter(ABC):
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def update_settings(self, settings: Mapping[str, Any]):
|
async def update_settings(self, settings: Mapping[str, Any]):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def filter(self, text: str) -> str:
|
async def filter(self, text: str) -> str:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def handle_interruption(self):
|
async def handle_interruption(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def reset_interruption(self):
|
async def reset_interruption(self):
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -33,12 +33,12 @@ class MarkdownTextFilter(BaseTextFilter):
|
|||||||
self._in_table = False
|
self._in_table = False
|
||||||
self._interrupted = False
|
self._interrupted = False
|
||||||
|
|
||||||
def update_settings(self, settings: Mapping[str, Any]):
|
async def update_settings(self, settings: Mapping[str, Any]):
|
||||||
for key, value in settings.items():
|
for key, value in settings.items():
|
||||||
if hasattr(self._settings, key):
|
if hasattr(self._settings, key):
|
||||||
setattr(self._settings, key, value)
|
setattr(self._settings, key, value)
|
||||||
|
|
||||||
def filter(self, text: str) -> str:
|
async def filter(self, text: str) -> str:
|
||||||
if self._settings.enable_text_filter:
|
if self._settings.enable_text_filter:
|
||||||
# Remove newlines and replace with a space only when there's no text before or after
|
# Remove newlines and replace with a space only when there's no text before or after
|
||||||
filtered_text = re.sub(r"^\s*\n", " ", text, flags=re.MULTILINE)
|
filtered_text = re.sub(r"^\s*\n", " ", text, flags=re.MULTILINE)
|
||||||
@@ -104,12 +104,12 @@ class MarkdownTextFilter(BaseTextFilter):
|
|||||||
else:
|
else:
|
||||||
return text
|
return text
|
||||||
|
|
||||||
def handle_interruption(self):
|
async def handle_interruption(self):
|
||||||
self._interrupted = True
|
self._interrupted = True
|
||||||
self._in_code_block = False
|
self._in_code_block = False
|
||||||
self._in_table = False
|
self._in_table = False
|
||||||
|
|
||||||
def reset_interruption(self):
|
async def reset_interruption(self):
|
||||||
self._interrupted = False
|
self._interrupted = False
|
||||||
|
|
||||||
#
|
#
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
Some inline code here
|
Some inline code here
|
||||||
"""
|
"""
|
||||||
|
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(result.strip(), expected_text.strip())
|
self.assertEqual(result.strip(), expected_text.strip())
|
||||||
|
|
||||||
async def test_space_preservation(self):
|
async def test_space_preservation(self):
|
||||||
@@ -45,7 +45,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
]
|
]
|
||||||
|
|
||||||
for text in input_text:
|
for text in input_text:
|
||||||
result = self.filter.filter(text)
|
result = await self.filter.filter(text)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
len(result), len(text), f"Space preservation failed for: '{text}'\nGot: '{result}'"
|
len(result), len(text), f"Space preservation failed for: '{text}'\nGot: '{result}'"
|
||||||
)
|
)
|
||||||
@@ -71,7 +71,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
for input_text, expected in test_cases.items():
|
for input_text, expected in test_cases.items():
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
result,
|
result,
|
||||||
expected,
|
expected,
|
||||||
@@ -88,7 +88,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
2. Second item
|
2. Second item
|
||||||
3. Third item with bold"""
|
3. Third item with bold"""
|
||||||
|
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
result.strip(),
|
result.strip(),
|
||||||
expected.strip(),
|
expected.strip(),
|
||||||
@@ -106,7 +106,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
for input_text, expected in test_cases.items():
|
for input_text, expected in test_cases.items():
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(result, expected, f"HTML entity conversion failed for: '{input_text}'")
|
self.assertEqual(result, expected, f"HTML entity conversion failed for: '{input_text}'")
|
||||||
|
|
||||||
async def test_asterisk_removal(self):
|
async def test_asterisk_removal(self):
|
||||||
@@ -120,7 +120,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
for input_text, expected in test_cases.items():
|
for input_text, expected in test_cases.items():
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(result, expected, f"Asterisk removal failed for: '{input_text}'")
|
self.assertEqual(result, expected, f"Asterisk removal failed for: '{input_text}'")
|
||||||
|
|
||||||
async def test_newline_handling(self):
|
async def test_newline_handling(self):
|
||||||
@@ -132,7 +132,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
for input_text, expected in test_cases.items():
|
for input_text, expected in test_cases.items():
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
result, expected, f"Newline handling failed for:\n{input_text}\nGot:\n{result}"
|
result, expected, f"Newline handling failed for:\n{input_text}\nGot:\n{result}"
|
||||||
)
|
)
|
||||||
@@ -148,7 +148,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
for input_text, expected in test_cases.items():
|
for input_text, expected in test_cases.items():
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
result,
|
result,
|
||||||
expected,
|
expected,
|
||||||
@@ -166,7 +166,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
for input_text, expected in test_cases.items():
|
for input_text, expected in test_cases.items():
|
||||||
result = self.filter.filter(input_text)
|
result = await self.filter.filter(input_text)
|
||||||
self.assertEqual(result, expected, f"Inline code handling failed for: '{input_text}'")
|
self.assertEqual(result, expected, f"Inline code handling failed for: '{input_text}'")
|
||||||
|
|
||||||
async def test_simple_table_removal(self):
|
async def test_simple_table_removal(self):
|
||||||
@@ -177,7 +177,7 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
expected = ""
|
expected = ""
|
||||||
|
|
||||||
result = filter.filter(input_text)
|
result = await filter.filter(input_text)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
result.strip(),
|
result.strip(),
|
||||||
expected.strip(),
|
expected.strip(),
|
||||||
@@ -198,15 +198,15 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
# Test with text filtering disabled
|
# Test with text filtering disabled
|
||||||
text_with_markdown = "**bold** and *italic* with `code`"
|
text_with_markdown = "**bold** and *italic* with `code`"
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
filter.filter(text_with_markdown),
|
await filter.filter(text_with_markdown),
|
||||||
text_with_markdown,
|
text_with_markdown,
|
||||||
"Disabled filter should not modify text",
|
"Disabled filter should not modify text",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Enable just text filtering
|
# Enable just text filtering
|
||||||
filter.update_settings({"enable_text_filter": True})
|
await filter.update_settings({"enable_text_filter": True})
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
filter.filter(text_with_markdown),
|
await filter.filter(text_with_markdown),
|
||||||
"bold and italic with code",
|
"bold and italic with code",
|
||||||
"Enabled filter should remove markdown",
|
"Enabled filter should remove markdown",
|
||||||
)
|
)
|
||||||
@@ -217,14 +217,18 @@ class TestMarkdownTextFilter(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
# Initial state - formatting should be removed
|
# Initial state - formatting should be removed
|
||||||
input_text = "**bold** and *italic*"
|
input_text = "**bold** and *italic*"
|
||||||
self.assertEqual(filter.filter(input_text), "bold and italic")
|
self.assertEqual(await filter.filter(input_text), "bold and italic")
|
||||||
|
|
||||||
# Disable text filtering
|
# Disable text filtering
|
||||||
filter.update_settings({"enable_text_filter": False})
|
await filter.update_settings({"enable_text_filter": False})
|
||||||
self.assertEqual(filter.filter(input_text), input_text, "Text filtering should be disabled")
|
self.assertEqual(
|
||||||
|
await filter.filter(input_text), input_text, "Text filtering should be disabled"
|
||||||
|
)
|
||||||
|
|
||||||
# Re-enable text filtering
|
# Re-enable text filtering
|
||||||
filter.update_settings({"enable_text_filter": True})
|
await filter.update_settings({"enable_text_filter": True})
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
filter.filter(input_text), "bold and italic", "Text filtering should be re-enabled"
|
await filter.filter(input_text),
|
||||||
|
"bold and italic",
|
||||||
|
"Text filtering should be re-enabled",
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user