File size: 4,507 Bytes
31dc8dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
from __future__ import annotations

from typing import Any, Callable

from diffulex.logger import get_logger
from diffulex.server.protocol import (
    ServingCommand,
    ServingEvent,
    ServingShutdown,
    serving_command_from_dict,
    serving_event_to_dict,
)
from diffulex.server.zmq_queue import ZmqPullQueue, ZmqPushQueue

logger = get_logger(__name__)


def default_engine_factory(model: str, **engine_kwargs):
    from diffulex import strategy as _strategy  # noqa: F401
    from diffulex.engine.engine import DiffulexEngine

    return DiffulexEngine(model, **engine_kwargs)


class SyncBackendWorker:
    def __init__(
        self,
        *,
        model: str,
        engine_kwargs: dict[str, Any] | None,
        recv_frontend,
        send_frontend,
        engine_factory: Callable[..., Any] | None = None,
        ready_queue=None,
    ) -> None:
        self.model = model
        self.engine_kwargs = dict(engine_kwargs or {})
        self.recv_frontend = recv_frontend
        self.send_frontend = send_frontend
        self.engine_factory = engine_factory or default_engine_factory
        self.ready_queue = ready_queue
        self.engine = None
        self.shutdown_requested = False

    @classmethod
    def from_zmq(
        cls,
        *,
        model: str,
        engine_kwargs: dict[str, Any] | None,
        command_addr: str,
        event_addr: str,
        ready_queue=None,
    ) -> "SyncBackendWorker":
        return cls(
            model=model,
            engine_kwargs=engine_kwargs,
            recv_frontend=ZmqPullQueue(command_addr, create=False, decoder=serving_command_from_dict),
            send_frontend=ZmqPushQueue(event_addr, create=False, encoder=serving_event_to_dict),
            ready_queue=ready_queue,
        )

    def init_engine(self) -> None:
        if self.engine is not None:
            return
        self.engine = self.engine_factory(self.model, **self.engine_kwargs)
        if self.ready_queue is not None:
            self.ready_queue.put("SyncBackendWorker is ready")

    def run_forever(self) -> None:
        self.init_engine()
        try:
            while not self.shutdown_requested:
                self.normal_loop()
        finally:
            self.shutdown()

    def normal_loop(self) -> None:
        assert self.engine is not None
        commands = self.receive_commands(blocking=self.engine.is_finished())
        commands = self.process_input_commands(commands)
        if commands or not self.engine.is_finished():
            events = self.engine.run_serving_tick(commands)
            self.send_result(events)

    def receive_commands(self, *, blocking: bool) -> list[ServingCommand]:
        commands: list[ServingCommand] = []
        if blocking:
            self.run_when_idle()
            commands.append(self.recv_frontend.get())
        while not self.recv_frontend.empty():
            commands.append(self.recv_frontend.get())
        return commands

    def process_input_commands(self, commands: list[ServingCommand]) -> list[ServingCommand]:
        input_commands: list[ServingCommand] = []
        for command in commands:
            if isinstance(command, ServingShutdown):
                self.shutdown_requested = True
            else:
                input_commands.append(command)
        return input_commands

    def send_result(self, events: list[ServingEvent]) -> None:
        for event in events:
            self.send_frontend.put(event)

    def run_when_idle(self) -> None:
        logger.info("SyncBackendWorker is idle, waiting for new requests...")

    def shutdown(self) -> None:
        if self.engine is not None:
            try:
                self.engine.exit()
            finally:
                self.engine = None
        for queue in (self.recv_frontend, self.send_frontend):
            stop = getattr(queue, "stop", None)
            if stop is not None:
                stop()


def run_sync_backend_worker(
    *,
    model: str,
    engine_kwargs: dict[str, Any] | None,
    command_addr: str,
    event_addr: str,
    ready_queue=None,
) -> None:
    try:
        worker = SyncBackendWorker.from_zmq(
            model=model,
            engine_kwargs=engine_kwargs,
            command_addr=command_addr,
            event_addr=event_addr,
            ready_queue=ready_queue,
        )
        worker.run_forever()
    except Exception as exc:
        if ready_queue is not None:
            ready_queue.put({"error": repr(exc)})
        raise