add openjtalk worker pkg

This commit is contained in:
kale4eat
2024-03-05 12:17:15 +09:00
parent bc1058270c
commit 1eb8fb4b08
5 changed files with 303 additions and 0 deletions

View File

@@ -0,0 +1,100 @@
"""
Run the pyopenjtalk worker in a separate process
to avoid user dictionary access error
"""
from typing import Optional, Any
from .worker_common import WOKER_PORT
from .worker_client import WorkerClient
from common.log import logger
WORKER_CLIENT: Optional[WorkerClient] = None
# pyopenjtalk interface
# g2p: not used
def run_frontend(text: str) -> list[dict[str, Any]]:
assert WORKER_CLIENT
ret = WORKER_CLIENT.dispatch_pyopenjtalk("run_frontend", text)
assert isinstance(ret, list)
return ret
def make_label(njd_features) -> list[str]:
assert WORKER_CLIENT
ret = WORKER_CLIENT.dispatch_pyopenjtalk("make_label", njd_features)
assert isinstance(ret, list)
return ret
def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None):
assert WORKER_CLIENT
WORKER_CLIENT.dispatch_pyopenjtalk("mecab_dict_index", path, out_path, dn_mecab)
def update_global_jtalk_with_user_dict(path: str):
assert WORKER_CLIENT
WORKER_CLIENT.dispatch_pyopenjtalk("update_global_jtalk_with_user_dict", path)
def unset_user_dict():
assert WORKER_CLIENT
WORKER_CLIENT.dispatch_pyopenjtalk("unset_user_dict")
# initialize module when imported
def initialize(port: int = WOKER_PORT):
import time
import socket
import sys
import atexit
global WORKER_CLIENT
logger.debug("initialize")
if WORKER_CLIENT:
return
client = None
try:
client = WorkerClient(port)
except (socket.timeout, socket.error):
logger.debug("try starting worker server")
import os
import subprocess
worker_pkg_path = os.path.relpath(
os.path.dirname(__file__), os.getcwd()
).replace(os.sep, ".")
subprocess.Popen([sys.executable, "-m", worker_pkg_path, "--port", str(port)])
# wait until server listening
count = 0
while True:
try:
client = WorkerClient(port)
break
except socket.error:
time.sleep(1)
count += 1
# 10: max number of retries
if count == 10:
raise TimeoutError("サーバーに接続できませんでした")
WORKER_CLIENT = client
def terminate():
global WORKER_CLIENT
if not WORKER_CLIENT:
return
if WORKER_CLIENT.status().get("client-count") == 1:
WORKER_CLIENT.quit_server()
WORKER_CLIENT.close()
WORKER_CLIENT = None
atexit.register(terminate)

View File

@@ -0,0 +1,16 @@
import argparse
from .worker_server import WorkerServer
from .worker_common import WOKER_PORT
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--port", type=int, default=WOKER_PORT)
args = parser.parse_args()
server = WorkerServer()
server.start_server(port=args.port)
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,40 @@
from typing import Any
import socket
from .worker_common import RequestType, receive_data, send_data
class WorkerClient:
def __init__(self, port: int) -> None:
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
# 5: timeout
sock.settimeout(5)
sock.connect((socket.gethostname(), port))
self.sock = sock
def __enter__(self):
return self
def __exit__(self, exc_type, exc_value, traceback):
self.close()
def close(self):
self.sock.close()
def dispatch_pyopenjtalk(self, func: str, *args, **kwargs):
data = {
"request-type": RequestType.PYOPENJTALK,
"func": func,
"args": args,
"kwargs": kwargs,
}
send_data(self.sock, data)
return receive_data(self.sock).get("return")
def status(self):
send_data(self.sock, {"request-type": RequestType.STATUS})
return receive_data(self.sock)
def quit_server(self):
send_data(self.sock, {"request-type": RequestType.QUIT_SERVER})
receive_data(self.sock)

View File

@@ -0,0 +1,44 @@
from typing import Any, Optional, Final
from enum import IntEnum, auto
import socket
import json
WOKER_PORT: Final[int] = 7861
HEADER_SIZE: Final[int] = 4
class RequestType(IntEnum):
STATUS = auto()
QUIT_SERVER = auto()
PYOPENJTALK = auto()
class ConnectionClosedException(Exception):
pass
# socket communication
def send_data(sock: socket.socket, data: dict[str, Any]):
json_data = json.dumps(data).encode()
header = len(json_data).to_bytes(HEADER_SIZE, byteorder="big")
sock.sendall(header + json_data)
def _receive_until(sock: socket.socket, size: int):
data = b""
while len(data) < size:
part = sock.recv(size - len(data))
if part == b"":
raise ConnectionClosedException("接続が閉じられました")
data += part
return data
def receive_data(sock: socket.socket) -> dict[str, Any]:
header = _receive_until(sock, HEADER_SIZE)
data_length = int.from_bytes(header, byteorder="big")
body = _receive_until(sock, data_length)
return json.loads(body.decode())

View File

@@ -0,0 +1,103 @@
import pyopenjtalk
import socket
import select
from .worker_common import (
ConnectionClosedException,
RequestType,
receive_data,
send_data,
)
from common.log import logger
# To make it as fast as possible
# Probably faster than calling getattr every time
_PYOPENJTALK_FUNC_DICT = {
"run_frontend": pyopenjtalk.run_frontend,
"make_label": pyopenjtalk.make_label,
"mecab_dict_index": pyopenjtalk.mecab_dict_index,
"update_global_jtalk_with_user_dict": pyopenjtalk.update_global_jtalk_with_user_dict,
"unset_user_dict": pyopenjtalk.unset_user_dict,
}
class WorkerServer:
def __init__(self) -> None:
self.client_count: int = 0
self.quit: bool = False
def handle_request(self, request):
request_type = None
try:
request_type = RequestType(request.get("request-type"))
except Exception:
return {
"success": False,
"reason": "request-type is invalid",
}
if request_type:
if request_type == RequestType.STATUS:
response = {
"success": True,
"client-count": self.client_count,
}
elif request_type == RequestType.QUIT_SERVER:
self.quit = True
response = {"success": True}
elif request_type == RequestType.PYOPENJTALK:
func_name = request.get("func")
assert isinstance(func_name, str)
func = _PYOPENJTALK_FUNC_DICT[func_name]
args = request.get("args")
kwargs = request.get("kwargs")
assert isinstance(args, list)
assert isinstance(kwargs, dict)
ret = func(*args, **kwargs)
response = {"success": True, "return": ret}
else:
# NOT REACHED
response = request
return response
def start_server(self, port: int):
logger.info("start pyopenjtalk worker server")
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:
server_socket.bind((socket.gethostname(), port))
server_socket.listen()
sockets = [server_socket]
while True:
ready_sockets, _, _ = select.select(sockets, [], [], 0.1)
for sock in ready_sockets:
if sock is server_socket:
logger.info("new client connected")
client_socket, _ = server_socket.accept()
sockets.append(client_socket)
self.client_count += 1
else:
# client
try:
request = receive_data(sock)
except ConnectionClosedException as e:
sock.close()
sockets.remove(sock)
self.client_count -= 1
logger.info("close connection")
continue
logger.debug(f"receive request: {request}")
response = self.handle_request(request)
logger.debug(f"send response: {response}")
try:
send_data(sock, response)
except Exception:
logger.warning(
"an exception occurred during sending responce"
)
if self.quit:
logger.info("quit pyopenjtalk worker server")
return