File size: 4,138 Bytes
b2c1c67
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
48954a5
 
 
b2c1c67
48954a5
 
 
b2c1c67
48954a5
 
 
 
 
 
b2c1c67
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
06f0a90
b2c1c67
 
 
 
 
 
 
 
 
 
06f0a90
b2c1c67
 
 
 
 
 
 
 
 
 
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
import os
import tempfile

_TMP_DIR = tempfile.mkdtemp(prefix="synapse-test-")
os.environ.update(
    {
        "SECRET_KEY": "test-secret-key-not-for-production",
        "DATABASE_URL": f"sqlite+aiosqlite:///{_TMP_DIR}/test.db",
        "VECTOR_STORE": "memory",
        "RERANK_ENABLED": "false",
        "OPENAI_API_KEY": "test-key",
        "CHAT_MODEL": "gpt-5.6-luna",
        "AVAILABLE_MODELS": "gpt-5.6-luna,gpt-4o-mini",
        "SUMMARIZE_AFTER_MESSAGES": "1000",
    }
)

import hashlib
from collections.abc import AsyncIterator
from typing import Any

import numpy as np
import pytest
from httpx import ASGITransport, AsyncClient

from app.api import deps
from app.database import Base, engine
from app.llm.base import ChatOptions, LLMProvider, ToolCallRequest
from app.llm.registry import set_provider
from app.main import app
from app.rag.vectorstore import MemoryVectorStore, set_vector_store


def deterministic_embedding(text: str, dims: int = 32) -> list[float]:
    vector = np.zeros(dims, dtype=np.float64)
    for token in text.lower().split():
        seed = int(hashlib.md5(token.encode()).hexdigest()[:8], 16)
        rng = np.random.RandomState(seed)
        vector += rng.randn(dims)
    norm = np.linalg.norm(vector)
    if norm > 0:
        vector /= norm
    return vector.tolist()


class FakeProvider(LLMProvider):
    def __init__(self, turns: list[dict[str, Any]] | None = None) -> None:
        self.turns = turns or [{"text": "Hello from the fake model."}]
        self.calls: list[list[dict[str, Any]]] = []
        self.embed_calls: list[list[str]] = []
        self.complete_response = "Fake title"

    async def stream_chat(
        self, messages: list[dict[str, Any]], options: ChatOptions
    ) -> AsyncIterator[dict[str, Any]]:
        self.calls.append(messages)
        turn = self.turns.pop(0) if self.turns else {"text": "(no script left)"}
        if "text" in turn:
            for word in turn["text"].split(" "):
                yield {"type": "delta", "text": word + " "}
        if "tool_calls" in turn:
            calls = []
            for i, spec in enumerate(turn["tool_calls"]):
                name, args, *rest = spec
                calls.append(
                    ToolCallRequest(
                        id=f"call_{i}",
                        name=name,
                        arguments=args,
                        extra=rest[0] if rest else {},
                        function_extra=rest[1] if len(rest) > 1 else {},
                    )
                )
            yield {"type": "tool_calls", "calls": calls}
        yield {"type": "usage", "input_tokens": 100, "output_tokens": 20}

    async def complete(
        self, messages: list[dict[str, Any]], model: str, temperature: float = 0.3
    ) -> str:
        return self.complete_response

    async def embed(self, texts: list[str]) -> list[list[float]]:
        self.embed_calls.append(list(texts))
        return [deterministic_embedding(text) for text in texts]


@pytest.fixture
async def fake_provider() -> AsyncIterator[FakeProvider]:
    provider = FakeProvider()
    set_provider(provider)
    yield provider
    set_provider(None)


@pytest.fixture
async def client(fake_provider: FakeProvider) -> AsyncIterator[AsyncClient]:
    await engine.dispose()
    async with engine.begin() as conn:
        await conn.run_sync(Base.metadata.drop_all)
        await conn.run_sync(Base.metadata.create_all)
    set_vector_store(MemoryVectorStore())
    deps.auth_limiter.reset()
    deps.chat_limiter.reset()
    transport = ASGITransport(app=app)
    async with AsyncClient(transport=transport, base_url="http://test") as http:
        yield http
    set_vector_store(None)
    await engine.dispose()


async def register_and_login(client: AsyncClient, email: str = "adi@example.com") -> dict[str, str]:
    response = await client.post(
        "/api/auth/register",
        json={"email": email, "username": "adi", "password": "supersecret123"},
    )
    assert response.status_code == 201, response.text
    tokens = response.json()
    return {"Authorization": f"Bearer {tokens['access_token']}"}