diff --git a/websocket/asr_server.py b/websocket/asr_server.py index d4a8410..578d748 100755 --- a/websocket/asr_server.py +++ b/websocket/asr_server.py @@ -9,6 +9,7 @@ import concurrent.futures import logging from vosk import Model, SpkModel, KaldiRecognizer +from asr_server_filter import Filter def process_chunk(rec, message): if message == '{"eof" : 1}': @@ -30,6 +31,8 @@ async def recognize(websocket, path): sample_rate = args.sample_rate show_words = args.show_words max_alternatives = args.max_alternatives + apply_filter = args.apply_filter + p_filter = None if not apply_filter else Filter() logging.info('Connection from %s', websocket.remote_address); @@ -63,11 +66,14 @@ async def recognize(websocket, path): rec.SetSpkModel(spk_model) response, stop = await loop.run_in_executor(pool, process_chunk, rec, message) + + if apply_filter: + response = p_filter.filter(response) + await websocket.send(response) if stop: break - def start(): global model @@ -92,6 +98,7 @@ def start(): args.sample_rate = float(os.environ.get('VOSK_SAMPLE_RATE', 8000)) args.max_alternatives = int(os.environ.get('VOSK_ALTERNATIVES', 0)) args.show_words = bool(os.environ.get('VOSK_SHOW_WORDS', True)) + args.apply_filter = bool(os.environ.get('VOSK_FILTER', True)) if len(sys.argv) > 1: args.model_path = sys.argv[1] diff --git a/websocket/asr_server_filter.py b/websocket/asr_server_filter.py new file mode 100644 index 0000000..aabf68a --- /dev/null +++ b/websocket/asr_server_filter.py @@ -0,0 +1,28 @@ +#!/usr/bin/env python3 + +import json +#import logging +from profanity_filter import ProfanityFilter +from profanity_check import predict + +class Filter: + + def __init__(self): + self.pf = ProfanityFilter() + + def filter(self, response: str): + py_json_response = self.apply_filter(json.loads(response)) + return json.dumps(py_json_response) + + def apply_filter(self, response: dict): + if "partial" in response: + text_type = "partial" + elif "text" in response: + text_type = "text" + transcript = response[text_type] + has_profanity = predict([transcript])[0] + #logging.info("Transcript is profane? %s", (transcript, has_profanity)) + if has_profanity: + censored_transcript = self.pf.censor(transcript) + response[text_type] = censored_transcript + return response