from dataclasses import dataclass from heapq import heappop, heappush, nlargest ROOT_NODE_OFFSET = 2 @dataclass(slots=True) class Candidate: text: str = "" is_complete: bool = False score: int = 0 address: int = 0 class OnDeviceHeadModel: def __init__(self, path): self.fp = open(path, "rb") self.address_size = self._read_int(1) self.score_size = self._read_int(1) if self.address_size not in (3, 4): raise ValueError(f"invalid address_size={self.address_size}") if self.score_size not in (2, 3, 4): raise ValueError(f"invalid score_size={self.score_size}") def _read_int(self, size): b = self.fp.read(size) if len(b) != size: raise EOFError return int.from_bytes(b, "little") def _node_info(self, address): self.fp.seek(address) block = self._read_int(self.score_size) leaf = None if block & 1: leaf = Candidate(is_complete=True, score=self._read_int(self.score_size)) return block >> 1, leaf def _child(self, prefix): b = self.fp.read(1) if not b: return None size = b[0] if size == 0: return None text = prefix + self.fp.read(size).decode("utf-8", errors="ignore") first = self.fp.read(1) if not first: return None first = first[0] if first & 1 == 0: score = int.from_bytes(bytes([first]) + self.fp.read(self.score_size - 1), "little") >> 1 return Candidate(text=text, is_complete=True, score=score) address = int.from_bytes(bytes([first]) + self.fp.read(self.address_size - 1), "little") >> 1 pos = self.fp.tell() score, _ = self._node_info(address) self.fp.seek(pos) return Candidate(text=text, score=score, address=address) def _children(self, node): if node.is_complete: return _, leaf = self._node_info(node.address) if leaf: leaf.text = node.text yield leaf while child := self._child(node.text): yield child def _find_start(self, prefix): node = Candidate(address=ROOT_NODE_OFFSET) while len(node.text) < len(prefix): for child in self._children(node): if child.text.startswith(prefix) or prefix.startswith(child.text): if child.is_complete and len(child.text) < len(prefix): continue node = child break else: return None return node def suggest(self, prefix, limit=10): prefix = prefix.strip().lower() if not prefix: return [] start = self._find_start(prefix) if not start: return [] heap = [(-start.score, start)] leaves = [] while heap: _, node = heappop(heap) for child in self._children(node): if child.is_complete: leaves.append(child) else: heappush(heap, (-child.score, child)) return [(c.text, c.score) for c in nlargest(limit, leaves, key=lambda c: c.score)]