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"