From 9bc7828e63a192c69dec7cfca1b4ab9f980cb0fb Mon Sep 17 00:00:00 2001 From: lhl Date: Sat, 29 Aug 2026 22:27:24 +0800 Subject: [PATCH] =?UTF-8?q?feat(rag):=20=E6=96=B0=E5=A2=9E=20Embedder=20?= =?UTF-8?q?=E6=8A=BD=E8=B1=A1=E4=B8=8E=E7=A6=BB=E7=BA=BF=20FakeEmbedder?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/genesis/rag/__init__.py | 0 src/genesis/rag/embeddings.py | 51 +++++++++++++++++++++++++++++++++++ tests/test_rag_embeddings.py | 30 +++++++++++++++++++++ 3 files changed, 81 insertions(+) create mode 100644 src/genesis/rag/__init__.py create mode 100644 src/genesis/rag/embeddings.py create mode 100644 tests/test_rag_embeddings.py diff --git a/src/genesis/rag/__init__.py b/src/genesis/rag/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/genesis/rag/embeddings.py b/src/genesis/rag/embeddings.py new file mode 100644 index 0000000..c8ba40e --- /dev/null +++ b/src/genesis/rag/embeddings.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +from typing import List, Protocol + + +class Embedder(Protocol): + def embed(self, texts: List[str]) -> List[List[float]]: ... + + +_DIM = 64 + + +def _tokenize(text: str) -> List[str]: + toks = [] + cur = "" + for ch in text.lower(): + if ch.isalnum(): + cur += ch + else: + if cur: + toks.append(cur) + cur = "" + if cur: + toks.append(cur) + out = [] + for t in toks: + idx = 0 + for i, c in enumerate(t): + if i > 0 and c.isupper(): + out.append(t[idx:i]) + idx = i + out.append(t[idx:]) + return [x for x in out if x] + + +class FakeEmbedder: + def embed(self, texts: List[str]) -> List[List[float]]: + vecs = [] + for t in texts: + v = [0.0] * _DIM + for tok in _tokenize(t): + h = __import__("hashlib").md5(tok.encode("utf-8")).digest() + idx = h[0] % _DIM + v[idx] += 1.0 + norm = __import__("math").sqrt(sum(x * x for x in v)) or 1.0 + vecs.append([x / norm for x in v]) + return vecs + + +def get_embedder(engine) -> Embedder: + return FakeEmbedder() diff --git a/tests/test_rag_embeddings.py b/tests/test_rag_embeddings.py new file mode 100644 index 0000000..b1f9770 --- /dev/null +++ b/tests/test_rag_embeddings.py @@ -0,0 +1,30 @@ +# tests/test_rag_embeddings.py +from genesis.rag.embeddings import FakeEmbedder, get_embedder + + +def test_fake_embedder_deterministic(): + e = FakeEmbedder() + a = e.embed(["OrderController 创建订单"])[0] + b = e.embed(["OrderController 创建订单"])[0] + assert a == b + + +def test_fake_embedder_similar_closer_than_unrelated(): + e = FakeEmbedder() + base = e.embed(["OrderController 处理创建订单请求"])[0] + sim = e.embed(["OrderController 保存订单到数据库"])[0] + dif = e.embed(["用户登录认证模块"])[0] + import math + + def cos(x, y): + dot = sum(p * q for p, q in zip(x, y)) + nx = math.sqrt(sum(p * p for p in x)) + ny = math.sqrt(sum(q * q for q in y)) + return dot / (nx * ny or 1.0) + + assert cos(base, sim) > cos(base, dif) + + +def test_get_embedder_fake_engine_returns_fake(): + e = get_embedder("fake") + assert e.__class__.__name__ == "FakeEmbedder"