add openjtalk worker pkg
This commit is contained in:
100
text/pyopenjtalk_worker/__init__.py
Normal file
100
text/pyopenjtalk_worker/__init__.py
Normal 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)
|
||||||
16
text/pyopenjtalk_worker/__main__.py
Normal file
16
text/pyopenjtalk_worker/__main__.py
Normal 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()
|
||||||
40
text/pyopenjtalk_worker/worker_client.py
Normal file
40
text/pyopenjtalk_worker/worker_client.py
Normal 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)
|
||||||
44
text/pyopenjtalk_worker/worker_common.py
Normal file
44
text/pyopenjtalk_worker/worker_common.py
Normal 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())
|
||||||
103
text/pyopenjtalk_worker/worker_server.py
Normal file
103
text/pyopenjtalk_worker/worker_server.py
Normal 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
|
||||||
Reference in New Issue
Block a user