Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
be31a2774f | ||
|
|
31f2c6451f | ||
|
|
e00a2db81b | ||
|
|
128cc5a603 | ||
|
|
475a6003f4 | ||
|
|
b862242418 | ||
|
|
41bad19a76 | ||
|
|
ed445eb07f | ||
|
|
053c1a6b11 | ||
|
|
1904f6391c | ||
|
|
5781484f23 | ||
|
|
2c484e71ab | ||
|
|
df8d292a6e | ||
|
|
7e76456599 | ||
|
|
a4ee545314 | ||
|
|
888abecd7f | ||
|
|
4940dfeffd | ||
|
|
c5999afcad | ||
|
|
3ad3dc517f | ||
|
|
0b77019f6c | ||
|
|
a4a436c020 | ||
|
|
1278d5c332 | ||
|
|
697e926dfe | ||
|
|
2ee246fad6 | ||
|
|
c6ff93c764 | ||
|
|
e11ca4b376 | ||
|
|
9bc1788ae2 | ||
|
|
5515714ae6 | ||
|
|
6f571b9e12 | ||
|
|
bab9c2ec08 | ||
|
|
63965696f4 | ||
|
|
7fb55c03fc | ||
|
|
bf0e3e133f |
@@ -0,0 +1,21 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
name: CI
|
||||
on:
|
||||
push:
|
||||
branches: [main, master]
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Python 3.12
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -e .
|
||||
- name: Run tests
|
||||
run: make verify
|
||||
|
||||
@@ -1,2 +1,6 @@
|
||||
__pycache__/
|
||||
*.pyc
|
||||
.env.json
|
||||
logs/
|
||||
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
# typosaurus-sandbox
|
||||
|
||||
A minimal Python calculator used to verify the Typosaurus agent system.
|
||||
@@ -7,8 +8,30 @@ A minimal Python calculator used to verify the Typosaurus agent system.
|
||||
- Full type annotations on every function signature.
|
||||
- No comments or docstrings in source files.
|
||||
|
||||
## Python backend
|
||||
|
||||
- Package manifest: `pyproject.toml`
|
||||
- Entry module: `src/typosaurus_sandbox/__main__.py`
|
||||
- Framework: FastAPI
|
||||
- Serve frontend: no
|
||||
|
||||
## Architecture
|
||||
|
||||
- Backend serves frontend: no
|
||||
- Module root: `src/typosaurus_sandbox/`
|
||||
- Calculator business logic: `src/typosaurus_sandbox/domain/calculator/operations.py`
|
||||
- HTTP API layer: `src/typosaurus_sandbox/presentation/api/v1/calculator.py`
|
||||
|
||||
## Verification
|
||||
|
||||
```
|
||||
make verify
|
||||
```
|
||||
|
||||
## CI
|
||||
|
||||
- Workflow file: `.gitea/workflows/ci.yml`
|
||||
- Trigger: push to `main` or `master` branches
|
||||
- Steps: checkout, Python 3.12 setup, dependency install, `make verify`
|
||||
|
||||
|
||||
|
||||
@@ -1,2 +1,8 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
verify:
|
||||
@python3 -m compileall -q src tests && python3 -m unittest discover -s tests -q && echo "verification passed"
|
||||
@PYTHONPATH=src python3 -m compileall -q src tests && PYTHONPATH=src python3 -m unittest discover -s tests -q && echo "verification passed"
|
||||
|
||||
run:
|
||||
@PYTHONPATH=src python3 -m typosaurus_sandbox
|
||||
|
||||
|
||||
@@ -1,3 +1,208 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
# typosaurus-sandbox
|
||||
|
||||
Sandbox for Typosaurus end-to-end verification
|
||||
Sandbox for Typosaurus end-to-end verification.
|
||||
|
||||
A FastAPI application serving arithmetic operations over HTTP with JSON request/response bodies.
|
||||
|
||||
## Configuration
|
||||
|
||||
The application uses a single `.env.json` file at the project root as its central point of truth
|
||||
for configuration. Defaults are plug-and-play and require no setup:
|
||||
|
||||
```json
|
||||
{
|
||||
"host": "127.0.0.1",
|
||||
"port": 8000
|
||||
}
|
||||
```
|
||||
|
||||
When no `.env.json` is present, the application starts with these defaults. To customise, create
|
||||
`.env.json` in the project root and populate only the keys that differ.
|
||||
|
||||
## Usage
|
||||
|
||||
### Start the server
|
||||
|
||||
```sh
|
||||
python -m typosaurus_sandbox
|
||||
```
|
||||
|
||||
The server listens on `http://127.0.0.1:8000` by default.
|
||||
|
||||
### Health check
|
||||
|
||||
```
|
||||
GET /health
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"status": "ok"}
|
||||
```
|
||||
|
||||
## API endpoints
|
||||
|
||||
All calculator endpoints accept `POST` requests with a JSON body and return a JSON response.
|
||||
|
||||
### POST /api/v1/calculator/add
|
||||
|
||||
Add two integers.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"left": 3, "right": 5}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 8}
|
||||
```
|
||||
|
||||
### POST /api/v1/calculator/subtract
|
||||
|
||||
Subtract the right integer from the left.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"left": 10, "right": 3}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 7}
|
||||
```
|
||||
|
||||
### POST /api/v1/calculator/clamp
|
||||
|
||||
Clamp a value between a low and high bound.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"value": 15, "low": 0, "high": 10}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 10}
|
||||
```
|
||||
|
||||
Boundaries are inclusive. A `low > high` combination produces a 422 validation response.
|
||||
|
||||
### POST /api/v1/calculator/clamp-to-byte
|
||||
|
||||
Clamp an integer to the byte range [0, 255].
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"value": 300}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 255}
|
||||
```
|
||||
|
||||
### POST /api/v1/calculator/average
|
||||
|
||||
Compute the arithmetic mean of a list of values.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"values": [1, 2, 3, 4, 5]}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 3.0}
|
||||
```
|
||||
|
||||
An empty list produces a 422 validation response.
|
||||
|
||||
### POST /api/v1/calculator/median
|
||||
|
||||
Compute the median of a list of values. Values are sorted internally; an even-length list returns the average of the two middle values as a float.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"values": [1, 3, 5]}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 3.0}
|
||||
```
|
||||
|
||||
Request (even length):
|
||||
|
||||
```json
|
||||
{"values": [1, 2, 3, 4]}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 2.5}
|
||||
```
|
||||
|
||||
An empty list produces a 422 validation response.
|
||||
|
||||
### POST /api/v1/calculator/variance
|
||||
|
||||
Compute the population variance of a list of values.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"values": [1, 2, 3, 4, 5]}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 2.0}
|
||||
```
|
||||
|
||||
An empty list produces a 422 validation response.
|
||||
|
||||
### POST /api/v1/calculator/percentage
|
||||
|
||||
Compute what percentage `value` is of `total`.
|
||||
|
||||
Request:
|
||||
|
||||
```json
|
||||
{"value": 50, "total": 100}
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{"result": 50.0}
|
||||
```
|
||||
|
||||
A zero `total` produces a 422 validation response.
|
||||
|
||||
## Verification
|
||||
|
||||
```sh
|
||||
make verify
|
||||
```
|
||||
|
||||
Runs compile-all checks against all source and test files, then executes the full test suite.
|
||||
Zero warnings are tolerated.
|
||||
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import os
|
||||
|
||||
from flask import Flask
|
||||
from flask import jsonify
|
||||
from flask import make_response
|
||||
from flask import request
|
||||
from flask.wrappers import Response
|
||||
|
||||
from typosaurus_sandbox.domain.calculator import add
|
||||
from typosaurus_sandbox.domain.calculator import clamp
|
||||
from typosaurus_sandbox.domain.calculator import clamp_to_byte
|
||||
from typosaurus_sandbox.domain.calculator import subtract
|
||||
|
||||
app = Flask(__name__)
|
||||
|
||||
|
||||
@app.route('/')
|
||||
def index() -> Response:
|
||||
index_path = os.path.join(os.path.dirname(__file__), 'index.html')
|
||||
with open(index_path) as f:
|
||||
return make_response(f.read(), 200, {'Content-Type': 'text/html'})
|
||||
|
||||
|
||||
@app.route('/add', methods=['GET'])
|
||||
def add_route() -> Response:
|
||||
try:
|
||||
left = int(request.args['left'])
|
||||
right = int(request.args['right'])
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return make_response(jsonify({'error': 'Invalid or missing parameters'}), 400)
|
||||
return make_response(jsonify({'result': add(left, right)}), 200)
|
||||
|
||||
|
||||
@app.route('/subtract', methods=['GET'])
|
||||
def subtract_route() -> Response:
|
||||
try:
|
||||
left = int(request.args['left'])
|
||||
right = int(request.args['right'])
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return make_response(jsonify({'error': 'Invalid or missing parameters'}), 400)
|
||||
return make_response(jsonify({'result': subtract(left, right)}), 200)
|
||||
|
||||
|
||||
@app.route('/clamp', methods=['GET'])
|
||||
def clamp_route() -> Response:
|
||||
try:
|
||||
value = int(request.args['value'])
|
||||
low = int(request.args['low'])
|
||||
high = int(request.args['high'])
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return make_response(jsonify({'error': 'Invalid or missing parameters'}), 400)
|
||||
try:
|
||||
result = clamp(value, low, high)
|
||||
except ValueError:
|
||||
return make_response(jsonify({'error': 'Invalid or missing parameters'}), 400)
|
||||
return make_response(jsonify({'result': result}), 200)
|
||||
|
||||
|
||||
@app.route('/clamp_to_byte', methods=['GET'])
|
||||
def clamp_to_byte_route() -> Response:
|
||||
try:
|
||||
value = int(request.args['value'])
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return make_response(jsonify({'error': 'Invalid or missing parameters'}), 400)
|
||||
return make_response(jsonify({'result': clamp_to_byte(value)}), 200)
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
<!-- retoor <retoor@molodetz.nl> -->
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>Calculator</title>
|
||||
</head>
|
||||
<body>
|
||||
<h1>Calculator</h1>
|
||||
|
||||
<fieldset>
|
||||
<legend>Add / Subtract</legend>
|
||||
<input type="number" id="left" placeholder="Left operand">
|
||||
<input type="number" id="right" placeholder="Right operand">
|
||||
<button onclick="calculate('add')">Add</button>
|
||||
<button onclick="calculate('subtract')">Subtract</button>
|
||||
</fieldset>
|
||||
|
||||
<fieldset>
|
||||
<legend>Clamp</legend>
|
||||
<input type="number" id="value" placeholder="Value">
|
||||
<input type="number" id="low" placeholder="Low">
|
||||
<input type="number" id="high" placeholder="High">
|
||||
<button onclick="calculate('clamp')">Clamp</button>
|
||||
</fieldset>
|
||||
|
||||
<fieldset>
|
||||
<legend>Clamp to Byte</legend>
|
||||
<input type="number" id="byte_value" placeholder="Value">
|
||||
<button onclick="calculate('clamp_to_byte')">Clamp to Byte</button>
|
||||
</fieldset>
|
||||
|
||||
<p id="output"></p>
|
||||
|
||||
<script>
|
||||
function calculate(operation) {
|
||||
const resultEl = document.getElementById('output');
|
||||
let url;
|
||||
if (operation === 'add' || operation === 'subtract') {
|
||||
const left = document.getElementById('left').value;
|
||||
const right = document.getElementById('right').value;
|
||||
url = '/' + operation + '?left=' + encodeURIComponent(left) + '&right=' + encodeURIComponent(right);
|
||||
} else if (operation === 'clamp') {
|
||||
const value = document.getElementById('value').value;
|
||||
const low = document.getElementById('low').value;
|
||||
const high = document.getElementById('high').value;
|
||||
url = '/clamp?value=' + encodeURIComponent(value) + '&low=' + encodeURIComponent(low) + '&high=' + encodeURIComponent(high);
|
||||
} else if (operation === 'clamp_to_byte') {
|
||||
const value = document.getElementById('byte_value').value;
|
||||
url = '/clamp_to_byte?value=' + encodeURIComponent(value);
|
||||
}
|
||||
fetch(url)
|
||||
.then(function(response) {
|
||||
return response.json().then(function(data) {
|
||||
if (!response.ok) {
|
||||
resultEl.textContent = 'Error: ' + (data.error || 'Unknown error');
|
||||
} else {
|
||||
resultEl.textContent = 'Result: ' + data.result;
|
||||
}
|
||||
});
|
||||
})
|
||||
.catch(function() {
|
||||
resultEl.textContent = 'Error: Network error';
|
||||
});
|
||||
}
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,16 @@
|
||||
[project]
|
||||
name = "typosaurus-sandbox"
|
||||
version = "0.1.0"
|
||||
description = "Sandbox for Typosaurus end-to-end verification"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"fastapi",
|
||||
"uvicorn[standard]",
|
||||
]
|
||||
|
||||
[build-system]
|
||||
requires = ["setuptools"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -0,0 +1,3 @@
|
||||
fastapi
|
||||
uvicorn[standard]
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from typosaurus_sandbox.app import App
|
||||
from typosaurus_sandbox.core import Config, setup_logging
|
||||
|
||||
__all__ = ["App", "Config", "setup_logging"]
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import logging
|
||||
|
||||
import uvicorn
|
||||
|
||||
from typosaurus_sandbox.app import App
|
||||
from typosaurus_sandbox.core import Config, setup_logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
setup_logging()
|
||||
config = Config.load()
|
||||
logger.info("starting server on %s:%d", config.host, config.port)
|
||||
uvicorn.run(App, host=config.host, port=config.port, log_level="info")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
from typosaurus_sandbox.presentation.api.v1.calculator import calculator_router
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
App = FastAPI(title="typosaurus-sandbox")
|
||||
|
||||
|
||||
@App.on_event("startup")
|
||||
def on_startup() -> None:
|
||||
logger.info("application startup complete")
|
||||
|
||||
|
||||
@App.get("/health")
|
||||
def health() -> dict[str, str]:
|
||||
logger.debug("health check requested")
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
App.include_router(calculator_router)
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from typosaurus_sandbox.core.config import Config
|
||||
from typosaurus_sandbox.core.logging import setup_logging
|
||||
|
||||
__all__ = ["Config", "setup_logging"]
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Config:
|
||||
host: str = "127.0.0.1"
|
||||
port: int = 8000
|
||||
|
||||
@classmethod
|
||||
def load(cls) -> "Config":
|
||||
config_path = Path(".env.json")
|
||||
if not config_path.exists():
|
||||
logger.info("no .env.json found, using defaults")
|
||||
return cls()
|
||||
with config_path.open() as f:
|
||||
data = json.load(f)
|
||||
host = data.get("host", cls.host)
|
||||
port = data.get("port", cls.port)
|
||||
logger.info("loaded config from .env.json: host=%s port=%s", host, port)
|
||||
return cls(host=host, port=port)
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import logging
|
||||
import logging.handlers
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def setup_logging() -> None:
|
||||
log_dir = Path("logs")
|
||||
log_dir.mkdir(exist_ok=True)
|
||||
|
||||
handler = logging.handlers.RotatingFileHandler(
|
||||
log_dir / "typosaurus-sandbox.log",
|
||||
maxBytes=10 * 1024 * 1024,
|
||||
backupCount=5,
|
||||
)
|
||||
handler.setFormatter(
|
||||
logging.Formatter("%(asctime)s [%(levelname)s] %(name)s: %(message)s")
|
||||
)
|
||||
|
||||
logging.basicConfig(level=logging.DEBUG, handlers=[handler])
|
||||
logging.getLogger(__name__).info("logging configured")
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
@@ -0,0 +1,6 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from typosaurus_sandbox.domain.calculator.operations import add, average, clamp, clamp_to_byte, median, percentage, subtract, variance
|
||||
|
||||
__all__ = ["add", "average", "clamp", "clamp_to_byte", "median", "percentage", "subtract", "variance"]
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from typing import Sequence
|
||||
from typing import Sequence, Union
|
||||
|
||||
|
||||
def add(left: int, right: int) -> int:
|
||||
@@ -25,6 +25,23 @@ def clamp_to_byte(value: int) -> int:
|
||||
return max(0, min(255, value))
|
||||
|
||||
|
||||
def average(values: list[int | float]) -> float:
|
||||
if not values:
|
||||
raise ValueError
|
||||
return sum(values) / len(values)
|
||||
|
||||
|
||||
def median(values: list[float]) -> float:
|
||||
if not values:
|
||||
raise ValueError
|
||||
sorted_values = sorted(values)
|
||||
n = len(sorted_values)
|
||||
mid = n // 2
|
||||
if n % 2 == 1:
|
||||
return sorted_values[mid]
|
||||
return (sorted_values[mid - 1] + sorted_values[mid]) / 2.0
|
||||
|
||||
|
||||
def variance(values: Sequence[float]) -> float:
|
||||
if not values:
|
||||
raise ValueError
|
||||
@@ -32,3 +49,8 @@ def variance(values: Sequence[float]) -> float:
|
||||
return sum((x - mean) ** 2 for x in values) / len(values)
|
||||
|
||||
|
||||
def percentage(value: Union[int, float], total: Union[int, float]) -> float:
|
||||
if total == 0:
|
||||
raise ValueError
|
||||
return (value / total) * 100
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
@@ -0,0 +1 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
@@ -0,0 +1 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
@@ -0,0 +1,117 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from typosaurus_sandbox.domain.calculator import add, average, clamp, clamp_to_byte, median, percentage, subtract, variance
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
calculator_router = APIRouter(prefix="/api/v1/calculator")
|
||||
|
||||
|
||||
class AddRequest(BaseModel):
|
||||
left: int
|
||||
right: int
|
||||
|
||||
|
||||
class SubtractRequest(BaseModel):
|
||||
left: int
|
||||
right: int
|
||||
|
||||
|
||||
class ClampRequest(BaseModel):
|
||||
value: float
|
||||
low: float
|
||||
high: float
|
||||
|
||||
|
||||
class ClampToByteRequest(BaseModel):
|
||||
value: int = Field(ge=-2147483648, le=2147483647)
|
||||
|
||||
|
||||
class ValuesRequest(BaseModel):
|
||||
values: list[float]
|
||||
|
||||
|
||||
class PercentageRequest(BaseModel):
|
||||
value: float
|
||||
total: float
|
||||
|
||||
|
||||
class IntResult(BaseModel):
|
||||
result: int
|
||||
|
||||
|
||||
class FloatResult(BaseModel):
|
||||
result: float
|
||||
|
||||
|
||||
@calculator_router.post("/add", response_model=IntResult)
|
||||
def calculate_add(body: AddRequest) -> IntResult:
|
||||
logger.debug("add %d + %d", body.left, body.right)
|
||||
return IntResult(result=add(body.left, body.right))
|
||||
|
||||
|
||||
@calculator_router.post("/subtract", response_model=IntResult)
|
||||
def calculate_subtract(body: SubtractRequest) -> IntResult:
|
||||
logger.debug("subtract %d - %d", body.left, body.right)
|
||||
return IntResult(result=subtract(body.left, body.right))
|
||||
|
||||
|
||||
@calculator_router.post("/clamp", response_model=FloatResult)
|
||||
def calculate_clamp(body: ClampRequest) -> FloatResult:
|
||||
try:
|
||||
result = clamp(body.value, body.low, body.high)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=422, detail="low must not exceed high")
|
||||
return FloatResult(result=result)
|
||||
|
||||
|
||||
@calculator_router.post("/clamp-to-byte", response_model=IntResult)
|
||||
def calculate_clamp_to_byte(body: ClampToByteRequest) -> IntResult:
|
||||
logger.debug("clamp-to-byte %d", body.value)
|
||||
return IntResult(result=clamp_to_byte(body.value))
|
||||
|
||||
|
||||
@calculator_router.post("/average", response_model=FloatResult)
|
||||
def calculate_average(body: ValuesRequest) -> FloatResult:
|
||||
logger.debug("average of %d values", len(body.values))
|
||||
try:
|
||||
result = average(body.values)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=422, detail="values list must not be empty")
|
||||
return FloatResult(result=result)
|
||||
|
||||
|
||||
@calculator_router.post("/median", response_model=FloatResult)
|
||||
def calculate_median(body: ValuesRequest) -> FloatResult:
|
||||
logger.debug("median of %d values", len(body.values))
|
||||
try:
|
||||
result = median(body.values)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=422, detail="values list must not be empty")
|
||||
return FloatResult(result=result)
|
||||
|
||||
|
||||
@calculator_router.post("/variance", response_model=FloatResult)
|
||||
def calculate_variance(body: ValuesRequest) -> FloatResult:
|
||||
logger.debug("variance of %d values", len(body.values))
|
||||
try:
|
||||
result = variance(body.values)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=422, detail="values list must not be empty")
|
||||
return FloatResult(result=result)
|
||||
|
||||
|
||||
@calculator_router.post("/percentage", response_model=FloatResult)
|
||||
def calculate_percentage(body: PercentageRequest) -> FloatResult:
|
||||
logger.debug("percentage %f of %f", body.value, body.total)
|
||||
try:
|
||||
result = percentage(body.value, body.total)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=422, detail="total must not be zero")
|
||||
return FloatResult(result=result)
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from typosaurus_sandbox.research.cache import TTLCache
|
||||
from typosaurus_sandbox.research.client import RsearchClient, RsearchError
|
||||
from typosaurus_sandbox.research.config import ResearchConfig
|
||||
from typosaurus_sandbox.research.envelopes import (
|
||||
ChatResponse,
|
||||
ChatUsage,
|
||||
DeepReport,
|
||||
DescribeResponse,
|
||||
SearchGrade,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
from typosaurus_sandbox.research.frontier import (
|
||||
DedupStats,
|
||||
QueryFrontier,
|
||||
fingerprint_text,
|
||||
normalize_url,
|
||||
query_variants_from_result,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ChatResponse",
|
||||
"ChatUsage",
|
||||
"DedupStats",
|
||||
"DeepReport",
|
||||
"DescribeResponse",
|
||||
"QueryFrontier",
|
||||
"RsearchClient",
|
||||
"RsearchError",
|
||||
"ResearchConfig",
|
||||
"SearchGrade",
|
||||
"SearchResponse",
|
||||
"SearchResult",
|
||||
"TTLCache",
|
||||
"fingerprint_text",
|
||||
"normalize_url",
|
||||
"query_variants_from_result",
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Generic, TypeVar
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheEntry(Generic[T]):
|
||||
value: T
|
||||
expires_at: float
|
||||
|
||||
|
||||
class TTLCache(Generic[T]):
|
||||
def __init__(self, name: str, ttl_seconds: float) -> None:
|
||||
self._name = name
|
||||
self._ttl_seconds = ttl_seconds
|
||||
self._entries: dict[str, CacheEntry[T]] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def get(self, key: str) -> T | None:
|
||||
with self._lock:
|
||||
entry = self._entries.get(key)
|
||||
if entry is None:
|
||||
logger.debug("cache %s miss key=%s", self._name, key)
|
||||
return None
|
||||
if time.monotonic() >= entry.expires_at:
|
||||
del self._entries[key]
|
||||
logger.debug("cache %s expired key=%s", self._name, key)
|
||||
return None
|
||||
logger.debug("cache %s hit key=%s", self._name, key)
|
||||
return entry.value
|
||||
|
||||
def set(self, key: str, value: T) -> None:
|
||||
with self._lock:
|
||||
self._entries[key] = CacheEntry(value=value, expires_at=time.monotonic() + self._ttl_seconds)
|
||||
logger.debug("cache %s set key=%s ttl=%.0fs", self._name, key, self._ttl_seconds)
|
||||
|
||||
def clear(self) -> None:
|
||||
with self._lock:
|
||||
count = len(self._entries)
|
||||
self._entries.clear()
|
||||
logger.debug("cache %s cleared %d entries", self._name, count)
|
||||
|
||||
@@ -0,0 +1,212 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from typing import Any
|
||||
|
||||
from typosaurus_sandbox.research.cache import TTLCache
|
||||
from typosaurus_sandbox.research.config import ResearchConfig
|
||||
from typosaurus_sandbox.research.envelopes import ChatResponse, DescribeResponse, SearchResponse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_ERROR_LENGTH = 200
|
||||
|
||||
|
||||
class RsearchError(RuntimeError):
|
||||
def __init__(self, message: str, status_code: int | None = None) -> None:
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
def _multipart_body(field_name: str, filename: str, mime_type: str, payload: bytes) -> tuple[bytes, str]:
|
||||
boundary = "----rsearch-" + secrets.token_hex(8)
|
||||
head = (
|
||||
f"--{boundary}\r\n".encode()
|
||||
+ f'Content-Disposition: form-data; name="{field_name}"; filename="{filename}"\r\n'.encode()
|
||||
+ f"Content-Type: {mime_type}\r\n\r\n".encode()
|
||||
)
|
||||
tail = b"\r\n--" + boundary.encode() + b"--\r\n"
|
||||
return head + payload + tail, f"multipart/form-data; boundary={boundary}"
|
||||
|
||||
|
||||
def _content_hash(image_bytes: bytes) -> str:
|
||||
return hashlib.sha256(image_bytes).hexdigest()
|
||||
|
||||
|
||||
class RsearchClient:
|
||||
def __init__(self, config: ResearchConfig | None = None) -> None:
|
||||
self._config = config if config is not None else ResearchConfig()
|
||||
self._search_cache = TTLCache[SearchResponse]("search", self._config.search_cache_ttl_seconds)
|
||||
self._content_cache = TTLCache[str]("content", self._config.content_cache_ttl_seconds)
|
||||
self._describe_cache = TTLCache[DescribeResponse]("describe", self._config.content_cache_ttl_seconds)
|
||||
|
||||
@property
|
||||
def config(self) -> ResearchConfig:
|
||||
return self._config
|
||||
|
||||
def get_cached_content(self, url: str) -> str | None:
|
||||
return self._content_cache.get(url)
|
||||
|
||||
async def search(
|
||||
self,
|
||||
query: str,
|
||||
*,
|
||||
source: str | None = None,
|
||||
count: int | None = None,
|
||||
content: bool = False,
|
||||
type: str | None = None,
|
||||
deep: bool = False,
|
||||
ai: bool = False,
|
||||
cache: bool = True,
|
||||
) -> SearchResponse:
|
||||
params: dict[str, str] = {"query": query}
|
||||
if source is not None:
|
||||
params["source"] = source
|
||||
if count is not None:
|
||||
params["count"] = str(count)
|
||||
if content:
|
||||
params["content"] = "true"
|
||||
if type is not None:
|
||||
params["type"] = type
|
||||
if deep:
|
||||
params["deep"] = "true"
|
||||
if ai:
|
||||
params["ai"] = "true"
|
||||
if not cache:
|
||||
params["cache"] = "false"
|
||||
key = urllib.parse.urlencode(sorted(params.items()))
|
||||
if cache:
|
||||
cached_response = self._search_cache.get(key)
|
||||
if cached_response is not None:
|
||||
return cached_response
|
||||
timeout = self._config.deep_timeout_seconds if deep else self._config.request_timeout_seconds
|
||||
status, data = await asyncio.to_thread(self._request, "GET", "/search", params, None, None, timeout)
|
||||
response = SearchResponse.from_dict(data)
|
||||
if cache:
|
||||
self._search_cache.set(key, response)
|
||||
if content:
|
||||
for result in response.results:
|
||||
if result.content:
|
||||
self._content_cache.set(result.url, result.content)
|
||||
logger.info(
|
||||
"search query=%r source=%s count=%s deep=%s ai=%s results=%d",
|
||||
query,
|
||||
response.source,
|
||||
response.count,
|
||||
deep,
|
||||
ai,
|
||||
len(response.results),
|
||||
)
|
||||
return response
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
system: str | None = None,
|
||||
json_mode: bool = False,
|
||||
cache: bool = True,
|
||||
) -> ChatResponse:
|
||||
payload: dict[str, Any] = {"prompt": prompt}
|
||||
if system is not None:
|
||||
payload["system"] = system
|
||||
if json_mode:
|
||||
payload["json"] = True
|
||||
if not cache:
|
||||
payload["cache"] = False
|
||||
body = json.dumps(payload).encode()
|
||||
headers = {"Content-Type": "application/json"}
|
||||
status, data = await asyncio.to_thread(self._request, "POST", "/chat", None, body, headers, None)
|
||||
response = ChatResponse.from_dict(data)
|
||||
logger.info("chat prompt=%r cached=%s", prompt, response.cached)
|
||||
return response
|
||||
|
||||
async def describe(self, url: str) -> DescribeResponse:
|
||||
key = f"url:{url}"
|
||||
cached = self._describe_cache.get(key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
status, data = await asyncio.to_thread(self._request, "GET", "/describe", {"url": url}, None, None, None)
|
||||
response = DescribeResponse.from_dict(data)
|
||||
self._describe_cache.set(key, response)
|
||||
logger.info("describe url=%s", url)
|
||||
return response
|
||||
|
||||
async def describe_upload(self, image_bytes: bytes, *, filename: str, mime_type: str) -> DescribeResponse:
|
||||
body, content_type = _multipart_body("file", filename, mime_type, image_bytes)
|
||||
headers = {"Content-Type": content_type}
|
||||
return await self._describe_post(image_bytes, body, headers)
|
||||
|
||||
async def describe_raw(self, image_bytes: bytes, *, mime_type: str) -> DescribeResponse:
|
||||
headers = {"Content-Type": mime_type}
|
||||
return await self._describe_post(image_bytes, image_bytes, headers)
|
||||
|
||||
async def _describe_post(self, image_bytes: bytes, body: bytes, headers: dict[str, str]) -> DescribeResponse:
|
||||
key = "hash:" + _content_hash(image_bytes)
|
||||
cached = self._describe_cache.get(key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
status, data = await asyncio.to_thread(self._request, "POST", "/describe", None, body, headers, None)
|
||||
response = DescribeResponse.from_dict(data)
|
||||
self._describe_cache.set(key, response)
|
||||
logger.info("describe post size=%d", len(image_bytes))
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _error_message(data: dict[str, Any]) -> str:
|
||||
error = data.get("error")
|
||||
if isinstance(error, str) and error:
|
||||
return error
|
||||
detail = data.get("detail")
|
||||
if isinstance(detail, str) and detail:
|
||||
return detail
|
||||
title = data.get("title")
|
||||
if isinstance(title, str) and title:
|
||||
return title
|
||||
return json.dumps(data)[:MAX_ERROR_LENGTH]
|
||||
|
||||
def _request(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
params: dict[str, str] | None = None,
|
||||
payload: bytes | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> tuple[int, dict[str, Any]]:
|
||||
timeout_seconds = timeout if timeout is not None else self._config.request_timeout_seconds
|
||||
base_url = self._config.base_url
|
||||
if base_url.endswith("/"):
|
||||
base_url = base_url[:-1]
|
||||
url = base_url + path
|
||||
if params:
|
||||
url = url + "?" + urllib.parse.urlencode(params)
|
||||
request = urllib.request.Request(url, data=payload, method=method, headers=headers or {})
|
||||
try:
|
||||
with urllib.request.urlopen(request, timeout=timeout_seconds) as response:
|
||||
status = response.status
|
||||
body = response.read()
|
||||
except urllib.error.HTTPError as exc:
|
||||
status = exc.code
|
||||
body = exc.read()
|
||||
except urllib.error.URLError as exc:
|
||||
raise RsearchError(f"connection failure for {method} {path}: {exc.reason}") from exc
|
||||
if not body:
|
||||
raise RsearchError(f"empty response for {method} {path}", status)
|
||||
try:
|
||||
data = json.loads(body)
|
||||
except (json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||
raise RsearchError(f"invalid JSON for {method} {path}: {exc}", status) from exc
|
||||
if not isinstance(data, dict):
|
||||
raise RsearchError(f"unexpected response shape for {method} {path}", status)
|
||||
if status >= 400 or data.get("success") is False:
|
||||
raise RsearchError(self._error_message(data), status)
|
||||
return status, data
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResearchConfig:
|
||||
base_url: str = "https://rsearch.app.molodetz.nl"
|
||||
request_timeout_seconds: float = 30.0
|
||||
deep_timeout_seconds: float = 180.0
|
||||
search_cache_ttl_seconds: float = 300.0
|
||||
content_cache_ttl_seconds: float = 86400.0
|
||||
max_concurrency: int = 8
|
||||
default_count: int = 10
|
||||
|
||||
@classmethod
|
||||
def load(cls) -> "ResearchConfig":
|
||||
config_path = Path(".env.json")
|
||||
if not config_path.exists():
|
||||
logger.info("no .env.json found, using default research config")
|
||||
return cls()
|
||||
with config_path.open() as f:
|
||||
data = json.load(f)
|
||||
research = data.get("research", {})
|
||||
logger.info("loaded research config from .env.json")
|
||||
return cls(
|
||||
base_url=research.get("base_url", cls.base_url),
|
||||
request_timeout_seconds=research.get("request_timeout_seconds", cls.request_timeout_seconds),
|
||||
deep_timeout_seconds=research.get("deep_timeout_seconds", cls.deep_timeout_seconds),
|
||||
search_cache_ttl_seconds=research.get("search_cache_ttl_seconds", cls.search_cache_ttl_seconds),
|
||||
content_cache_ttl_seconds=research.get("content_cache_ttl_seconds", cls.content_cache_ttl_seconds),
|
||||
max_concurrency=research.get("max_concurrency", cls.max_concurrency),
|
||||
default_count=research.get("default_count", cls.default_count),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,203 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _as_float(value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchGrade:
|
||||
overall: float = 0.0
|
||||
relevance: float = 0.0
|
||||
depth: float = 0.0
|
||||
authority: float = 0.0
|
||||
freshness: float = 0.0
|
||||
word_count: int = 0
|
||||
intent_hits: int = 0
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any] | None) -> "SearchGrade | None":
|
||||
if data is None:
|
||||
return None
|
||||
return cls(
|
||||
overall=float(data.get("overall", 0.0) or 0.0),
|
||||
relevance=float(data.get("relevance", 0.0) or 0.0),
|
||||
depth=float(data.get("depth", 0.0) or 0.0),
|
||||
authority=float(data.get("authority", 0.0) or 0.0),
|
||||
freshness=float(data.get("freshness", 0.0) or 0.0),
|
||||
word_count=int(data.get("word_count", 0) or 0),
|
||||
intent_hits=int(data.get("intent_hits", 0) or 0),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchResult:
|
||||
title: str = ""
|
||||
url: str = ""
|
||||
description: str = ""
|
||||
source: str = ""
|
||||
content: str | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
index: int | None = None
|
||||
grade: SearchGrade | None = None
|
||||
query_origin: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "SearchResult":
|
||||
return cls(
|
||||
title=data.get("title", ""),
|
||||
url=data.get("url", ""),
|
||||
description=data.get("description", ""),
|
||||
source=data.get("source", ""),
|
||||
content=data.get("content"),
|
||||
extra=data.get("extra", {}),
|
||||
index=data.get("index"),
|
||||
grade=SearchGrade.from_dict(data.get("grade")),
|
||||
query_origin=data.get("query_origin"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DeepReport:
|
||||
query: str = ""
|
||||
markdown: str = ""
|
||||
sources: list[SearchResult] = field(default_factory=list)
|
||||
graded_count: int = 0
|
||||
total_count: int = 0
|
||||
model: str = ""
|
||||
elapsed: float = 0.0
|
||||
cache_hit: bool = False
|
||||
rounds: int = 0
|
||||
queries_tried: list[str] = field(default_factory=list)
|
||||
error: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any] | None) -> "DeepReport | None":
|
||||
if data is None:
|
||||
return None
|
||||
sources = [SearchResult.from_dict(item) for item in data.get("sources", [])]
|
||||
return cls(
|
||||
query=data.get("query", ""),
|
||||
markdown=data.get("markdown", ""),
|
||||
sources=sources,
|
||||
graded_count=int(data.get("graded_count", 0) or 0),
|
||||
total_count=int(data.get("total_count", 0) or 0),
|
||||
model=data.get("model", ""),
|
||||
elapsed=_as_float(data.get("elapsed")) or 0.0,
|
||||
cache_hit=bool(data.get("cache_hit", False)),
|
||||
rounds=int(data.get("rounds", 0) or 0),
|
||||
queries_tried=list(data.get("queries_tried", [])),
|
||||
error=data.get("error"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchResponse:
|
||||
query: str = ""
|
||||
source: str = ""
|
||||
count: int = 0
|
||||
results: list[SearchResult] = field(default_factory=list)
|
||||
success: bool = False
|
||||
error: str | None = None
|
||||
ai_response: str | None = None
|
||||
ai_error: str | None = None
|
||||
deep: DeepReport | None = None
|
||||
timestamp: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "SearchResponse":
|
||||
results = [SearchResult.from_dict(item) for item in data.get("results", [])]
|
||||
return cls(
|
||||
query=data.get("query", ""),
|
||||
source=data.get("source", ""),
|
||||
count=int(data.get("count", 0) or 0),
|
||||
results=results,
|
||||
success=bool(data.get("success", False)),
|
||||
error=data.get("error"),
|
||||
ai_response=data.get("ai_response"),
|
||||
ai_error=data.get("ai_error"),
|
||||
deep=DeepReport.from_dict(data.get("deep")),
|
||||
timestamp=data.get("timestamp"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatUsage:
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
cost_usd: float = 0.0
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any] | None) -> "ChatUsage | None":
|
||||
if data is None:
|
||||
return None
|
||||
return cls(
|
||||
prompt_tokens=int(data.get("prompt_tokens", 0) or 0),
|
||||
completion_tokens=int(data.get("completion_tokens", 0) or 0),
|
||||
total_tokens=int(data.get("total_tokens", 0) or 0),
|
||||
cost_usd=float(data.get("cost_usd", 0.0) or 0.0),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatResponse:
|
||||
response: str = ""
|
||||
prompt: str = ""
|
||||
json_mode: bool = False
|
||||
cached: bool = False
|
||||
usage: ChatUsage | None = None
|
||||
error: str | None = None
|
||||
max_context_window: int | None = None
|
||||
max_output_tokens: int | None = None
|
||||
elapsed: float | None = None
|
||||
timestamp: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "ChatResponse":
|
||||
return cls(
|
||||
response=data.get("response", ""),
|
||||
prompt=data.get("prompt", ""),
|
||||
json_mode=bool(data.get("json_mode", False)),
|
||||
cached=bool(data.get("cached", False)),
|
||||
usage=ChatUsage.from_dict(data.get("usage")),
|
||||
error=data.get("error"),
|
||||
max_context_window=data.get("max_context_window"),
|
||||
max_output_tokens=data.get("max_output_tokens"),
|
||||
elapsed=_as_float(data.get("elapsed")),
|
||||
timestamp=data.get("timestamp"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DescribeResponse:
|
||||
description: str = ""
|
||||
url: str | None = None
|
||||
mime_type: str | None = None
|
||||
size: int | None = None
|
||||
elapsed: float | None = None
|
||||
timestamp: str | None = None
|
||||
success: bool = True
|
||||
error: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "DescribeResponse":
|
||||
return cls(
|
||||
description=data.get("description", ""),
|
||||
url=data.get("url"),
|
||||
mime_type=data.get("mime_type"),
|
||||
size=data.get("size"),
|
||||
elapsed=_as_float(data.get("elapsed")),
|
||||
timestamp=data.get("timestamp"),
|
||||
success=bool(data.get("success", True)),
|
||||
error=data.get("error"),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
import re
|
||||
import threading
|
||||
import urllib.parse
|
||||
from dataclasses import dataclass
|
||||
|
||||
from typosaurus_sandbox.research.envelopes import SearchResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MIN_QUERY_LENGTH = 2
|
||||
MAX_QUERY_LENGTH = 200
|
||||
DEFAULT_PORTS: dict[str, int] = {"http": 80, "https": 443}
|
||||
|
||||
|
||||
def _clean_text(value: str) -> str:
|
||||
return " ".join(value.split())
|
||||
|
||||
|
||||
def _query_key(query: str) -> str:
|
||||
return _clean_text(query).casefold()
|
||||
|
||||
|
||||
def normalize_url(url: str) -> str:
|
||||
cleaned = _clean_text(url)
|
||||
try:
|
||||
parsed = urllib.parse.urlsplit(cleaned)
|
||||
except ValueError:
|
||||
return cleaned
|
||||
scheme = parsed.scheme.lower()
|
||||
if scheme not in DEFAULT_PORTS:
|
||||
return cleaned
|
||||
host = (parsed.hostname or "").lower()
|
||||
if not host:
|
||||
return cleaned
|
||||
try:
|
||||
host = host.encode("idna").decode("ascii")
|
||||
except UnicodeError:
|
||||
pass
|
||||
port: int | None = None
|
||||
try:
|
||||
port = parsed.port
|
||||
except ValueError:
|
||||
port = None
|
||||
if port is not None and DEFAULT_PORTS.get(scheme) == port:
|
||||
port = None
|
||||
display_host = f"[{host}]" if ":" in host else host
|
||||
netloc = display_host if port is None else f"{display_host}:{port}"
|
||||
path = re.sub(r"/{2,}", "/", parsed.path)
|
||||
if len(path) > 1 and path.endswith("/"):
|
||||
path = path[:-1]
|
||||
if parsed.query:
|
||||
return f"{scheme}://{netloc}{path}?{parsed.query}"
|
||||
return f"{scheme}://{netloc}{path}"
|
||||
|
||||
|
||||
def fingerprint_text(text: str) -> str:
|
||||
normalized = _clean_text(text)
|
||||
return hashlib.sha256(normalized.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def query_variants_from_result(result: SearchResult) -> list[tuple[str, str]]:
|
||||
variants: list[tuple[str, str]] = []
|
||||
if result.title:
|
||||
variants.append((result.title, "title"))
|
||||
if result.description:
|
||||
variants.append((result.description, "description"))
|
||||
for value in result.extra.values():
|
||||
if isinstance(value, str) and value:
|
||||
variants.append((value, "extra"))
|
||||
return variants
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DedupStats:
|
||||
queries_generated: int = 0
|
||||
queries_enqueued: int = 0
|
||||
queries_issued: int = 0
|
||||
queries_duplicates_skipped: int = 0
|
||||
urls_seen: int = 0
|
||||
urls_duplicates_skipped: int = 0
|
||||
content_seen: int = 0
|
||||
content_duplicates_skipped: int = 0
|
||||
|
||||
def to_dict(self) -> dict[str, int]:
|
||||
return {
|
||||
"queries_generated": self.queries_generated,
|
||||
"queries_enqueued": self.queries_enqueued,
|
||||
"queries_issued": self.queries_issued,
|
||||
"queries_duplicates_skipped": self.queries_duplicates_skipped,
|
||||
"urls_seen": self.urls_seen,
|
||||
"urls_duplicates_skipped": self.urls_duplicates_skipped,
|
||||
"content_seen": self.content_seen,
|
||||
"content_duplicates_skipped": self.content_duplicates_skipped,
|
||||
}
|
||||
|
||||
|
||||
class QueryFrontier:
|
||||
def __init__(self, subject: str | None = None) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._seen_queries: set[str] = set()
|
||||
self._seen_urls: set[str] = set()
|
||||
self._seen_content: set[str] = set()
|
||||
self._origins: dict[str, str] = {}
|
||||
self._pending: asyncio.Queue[str] = asyncio.Queue()
|
||||
self._queries_generated = 0
|
||||
self._queries_enqueued = 0
|
||||
self._queries_issued = 0
|
||||
self._queries_duplicates_skipped = 0
|
||||
self._urls_seen = 0
|
||||
self._urls_duplicates_skipped = 0
|
||||
self._content_seen = 0
|
||||
self._content_duplicates_skipped = 0
|
||||
if subject:
|
||||
self.seed(subject)
|
||||
|
||||
def seed(self, subject: str) -> None:
|
||||
cleaned = _clean_text(subject)
|
||||
if cleaned:
|
||||
self.push_query(cleaned, "seed")
|
||||
logger.info("frontier seeded subject=%r", cleaned)
|
||||
|
||||
def push_query(self, query: str, origin: str = "manual") -> bool:
|
||||
cleaned = _clean_text(query)
|
||||
if not MIN_QUERY_LENGTH <= len(cleaned) <= MAX_QUERY_LENGTH:
|
||||
logger.debug("query variant invalid length=%d query=%r", len(cleaned), cleaned)
|
||||
return False
|
||||
key = _query_key(cleaned)
|
||||
with self._lock:
|
||||
self._queries_generated += 1
|
||||
if key in self._seen_queries:
|
||||
self._queries_duplicates_skipped += 1
|
||||
logger.debug("query duplicate skipped origin=%s query=%r", origin, cleaned)
|
||||
return False
|
||||
self._seen_queries.add(key)
|
||||
self._origins[key] = origin
|
||||
self._queries_enqueued += 1
|
||||
self._pending.put_nowait(cleaned)
|
||||
logger.info("query enqueued origin=%s query=%r", origin, cleaned)
|
||||
return True
|
||||
|
||||
def push_variants_from_result(self, result: SearchResult) -> int:
|
||||
new_queries = 0
|
||||
for text, origin in query_variants_from_result(result):
|
||||
if self.push_query(text, origin):
|
||||
new_queries += 1
|
||||
return new_queries
|
||||
|
||||
def register_url(self, url: str) -> bool:
|
||||
if not url:
|
||||
return False
|
||||
normalized = normalize_url(url)
|
||||
with self._lock:
|
||||
if normalized in self._seen_urls:
|
||||
self._urls_duplicates_skipped += 1
|
||||
logger.debug("url duplicate skipped url=%s", normalized)
|
||||
return False
|
||||
self._seen_urls.add(normalized)
|
||||
self._urls_seen += 1
|
||||
logger.info("url registered url=%s", normalized)
|
||||
return True
|
||||
|
||||
def register_content(self, text: str) -> bool:
|
||||
if not text.strip():
|
||||
return False
|
||||
fingerprint = fingerprint_text(text)
|
||||
with self._lock:
|
||||
if fingerprint in self._seen_content:
|
||||
self._content_duplicates_skipped += 1
|
||||
logger.debug("content duplicate skipped fingerprint=%s", fingerprint)
|
||||
return False
|
||||
self._seen_content.add(fingerprint)
|
||||
self._content_seen += 1
|
||||
logger.info("content registered fingerprint=%s", fingerprint)
|
||||
return True
|
||||
|
||||
def register_result(self, result: SearchResult) -> bool:
|
||||
is_new = self.register_url(result.url)
|
||||
if result.content:
|
||||
self.register_content(result.content)
|
||||
return is_new
|
||||
|
||||
async def get_query(self) -> str:
|
||||
query = await self._pending.get()
|
||||
with self._lock:
|
||||
self._queries_issued += 1
|
||||
logger.info("query issued query=%r", query)
|
||||
return query
|
||||
|
||||
def pop_query(self) -> str | None:
|
||||
try:
|
||||
query = self._pending.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
return None
|
||||
with self._lock:
|
||||
self._queries_issued += 1
|
||||
logger.info("query issued query=%r", query)
|
||||
return query
|
||||
|
||||
def pending_count(self) -> int:
|
||||
return self._pending.qsize()
|
||||
|
||||
def has_pending(self) -> bool:
|
||||
return not self._pending.empty()
|
||||
|
||||
def origin_of(self, query: str) -> str | None:
|
||||
with self._lock:
|
||||
return self._origins.get(_query_key(query))
|
||||
|
||||
def snapshot(self) -> DedupStats:
|
||||
with self._lock:
|
||||
return DedupStats(
|
||||
queries_generated=self._queries_generated,
|
||||
queries_enqueued=self._queries_enqueued,
|
||||
queries_issued=self._queries_issued,
|
||||
queries_duplicates_skipped=self._queries_duplicates_skipped,
|
||||
urls_seen=self._urls_seen,
|
||||
urls_duplicates_skipped=self._urls_duplicates_skipped,
|
||||
content_seen=self._content_seen,
|
||||
content_duplicates_skipped=self._content_duplicates_skipped,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import unittest
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning, module="starlette")
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from typosaurus_sandbox.app import App
|
||||
|
||||
client = TestClient(App)
|
||||
|
||||
|
||||
class TestHealthEndpoint(unittest.TestCase):
|
||||
|
||||
def test_health_returns_ok(self) -> None:
|
||||
response = client.get("/health")
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"status": "ok"})
|
||||
|
||||
|
||||
class TestCalculatorAddEndpoint(unittest.TestCase):
|
||||
|
||||
def test_add_positive_integers(self) -> None:
|
||||
response = client.post("/api/v1/calculator/add", json={"left": 3, "right": 5})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 8})
|
||||
|
||||
def test_add_negative_integers(self) -> None:
|
||||
response = client.post("/api/v1/calculator/add", json={"left": -3, "right": -5})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": -8})
|
||||
|
||||
def test_add_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/add", json={"left": "abc", "right": 5})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_add_missing_field_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/add", json={"left": 3})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorSubtractEndpoint(unittest.TestCase):
|
||||
|
||||
def test_subtract_positive(self) -> None:
|
||||
response = client.post("/api/v1/calculator/subtract", json={"left": 10, "right": 3})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 7})
|
||||
|
||||
def test_subtract_negative_result(self) -> None:
|
||||
response = client.post("/api/v1/calculator/subtract", json={"left": 3, "right": 10})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": -7})
|
||||
|
||||
def test_subtract_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/subtract", json={"left": 10, "right": None})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_subtract_missing_left_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/subtract", json={"right": 3})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_subtract_missing_right_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/subtract", json={"left": 10})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorClampEndpoint(unittest.TestCase):
|
||||
|
||||
def test_clamp_value_below_low(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": -5, "low": 0, "high": 10})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 0})
|
||||
|
||||
def test_clamp_value_above_high(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": 15, "low": 0, "high": 10})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 10})
|
||||
|
||||
def test_clamp_value_in_range(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": 5, "low": 0, "high": 10})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 5.0})
|
||||
|
||||
def test_clamp_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": "x", "low": 0, "high": 10})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_clamp_low_greater_than_high_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": 5, "low": 10, "high": 0})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_clamp_missing_low_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": 5, "high": 10})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_clamp_missing_high_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp", json={"value": 5, "low": 0})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorClampToByteEndpoint(unittest.TestCase):
|
||||
|
||||
def test_clamp_to_byte_within_range(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp-to-byte", json={"value": 128})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 128})
|
||||
|
||||
def test_clamp_to_byte_below_zero(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp-to-byte", json={"value": -10})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 0})
|
||||
|
||||
def test_clamp_to_byte_above_255(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp-to-byte", json={"value": 300})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 255})
|
||||
|
||||
def test_clamp_to_byte_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/clamp-to-byte", json={"value": "abc"})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorAverageEndpoint(unittest.TestCase):
|
||||
|
||||
def test_average_positive_values(self) -> None:
|
||||
response = client.post("/api/v1/calculator/average", json={"values": [1, 2, 3, 4, 5]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 3.0})
|
||||
|
||||
def test_average_single_value(self) -> None:
|
||||
response = client.post("/api/v1/calculator/average", json={"values": [5]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 5.0})
|
||||
|
||||
def test_average_negative_values(self) -> None:
|
||||
response = client.post("/api/v1/calculator/average", json={"values": [-10, -20, -30]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": -20.0})
|
||||
|
||||
def test_average_empty_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/average", json={"values": []})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_average_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/average", json={"values": ["a", "b"]})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_average_missing_field_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/average", json={})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorMedianEndpoint(unittest.TestCase):
|
||||
|
||||
def test_median_odd_length(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={"values": [1, 3, 5]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 3.0})
|
||||
|
||||
def test_median_even_length(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={"values": [1, 2, 3, 4]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 2.5})
|
||||
|
||||
def test_median_single_element(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={"values": [7]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 7.0})
|
||||
|
||||
def test_median_unsorted_input(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={"values": [3, 1, 2]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 2.0})
|
||||
|
||||
def test_median_empty_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={"values": []})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_median_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={"values": ["a"]})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_median_missing_field_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/median", json={})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorVarianceEndpoint(unittest.TestCase):
|
||||
|
||||
def test_variance_known_set(self) -> None:
|
||||
response = client.post("/api/v1/calculator/variance", json={"values": [1, 2, 3, 4, 5]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 2.0})
|
||||
|
||||
def test_variance_constant_values(self) -> None:
|
||||
response = client.post("/api/v1/calculator/variance", json={"values": [1.0, 1.0, 1.0]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 0.0})
|
||||
|
||||
def test_variance_single_element(self) -> None:
|
||||
response = client.post("/api/v1/calculator/variance", json={"values": [42.0]})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 0.0})
|
||||
|
||||
def test_variance_empty_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/variance", json={"values": []})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_variance_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/variance", json={"values": ["a", "b", "c"]})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_variance_missing_field_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/variance", json={})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
class TestCalculatorPercentageEndpoint(unittest.TestCase):
|
||||
|
||||
def test_percentage_half(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"value": 50, "total": 100})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 50.0})
|
||||
|
||||
def test_percentage_quarter(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"value": 25, "total": 100})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 25.0})
|
||||
|
||||
def test_percentage_zero_value(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"value": 0, "total": 100})
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.json(), {"result": 0.0})
|
||||
|
||||
def test_percentage_total_zero_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"value": 50, "total": 0})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_percentage_invalid_input_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"value": "abc", "total": 100})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_percentage_missing_value_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"total": 100})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
def test_percentage_missing_total_returns_422(self) -> None:
|
||||
response = client.post("/api/v1/calculator/percentage", json={"value": 50})
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
|
||||
+131
-1
@@ -3,7 +3,69 @@
|
||||
import math
|
||||
import unittest
|
||||
|
||||
from src.calculator import clamp, variance
|
||||
from typosaurus_sandbox.domain.calculator import add, average, clamp, clamp_to_byte, median, percentage, subtract, variance
|
||||
|
||||
|
||||
class TestAddFunction(unittest.TestCase):
|
||||
|
||||
def test_add_positive_integers(self) -> None:
|
||||
self.assertEqual(add(3, 5), 8)
|
||||
|
||||
def test_add_negative_integers(self) -> None:
|
||||
self.assertEqual(add(-3, -5), -8)
|
||||
|
||||
def test_add_mixed_sign(self) -> None:
|
||||
self.assertEqual(add(-3, 5), 2)
|
||||
|
||||
|
||||
class TestAverageFunction(unittest.TestCase):
|
||||
|
||||
def test_empty_sequence_raises_value_error(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
average([])
|
||||
|
||||
def test_single_element(self) -> None:
|
||||
self.assertEqual(average([5]), 5.0)
|
||||
|
||||
def test_positive_values(self) -> None:
|
||||
self.assertEqual(average([1, 2, 3, 4, 5]), 3.0)
|
||||
|
||||
def test_negative_values(self) -> None:
|
||||
self.assertEqual(average([-10, -20, -30]), -20.0)
|
||||
|
||||
def test_mixed_positive_and_negative(self) -> None:
|
||||
self.assertEqual(average([-5, 0, 5]), 0.0)
|
||||
|
||||
def test_float_values(self) -> None:
|
||||
self.assertEqual(average([1.5, 2.5, 3.0]), 7.0 / 3.0)
|
||||
|
||||
|
||||
class TestSubtractFunction(unittest.TestCase):
|
||||
|
||||
def test_subtract_positive(self) -> None:
|
||||
self.assertEqual(subtract(10, 3), 7)
|
||||
|
||||
def test_subtract_negative_result(self) -> None:
|
||||
self.assertEqual(subtract(3, 10), -7)
|
||||
|
||||
def test_subtract_negative_numbers(self) -> None:
|
||||
self.assertEqual(subtract(-5, -3), -2)
|
||||
|
||||
|
||||
class TestClampToByteFunction(unittest.TestCase):
|
||||
|
||||
def test_clamp_to_byte_within_range(self) -> None:
|
||||
self.assertEqual(clamp_to_byte(128), 128)
|
||||
|
||||
def test_clamp_to_byte_below_zero(self) -> None:
|
||||
self.assertEqual(clamp_to_byte(-10), 0)
|
||||
|
||||
def test_clamp_to_byte_above_255(self) -> None:
|
||||
self.assertEqual(clamp_to_byte(300), 255)
|
||||
|
||||
def test_clamp_to_byte_boundaries(self) -> None:
|
||||
self.assertEqual(clamp_to_byte(0), 0)
|
||||
self.assertEqual(clamp_to_byte(255), 255)
|
||||
|
||||
|
||||
class TestClampFunction(unittest.TestCase):
|
||||
@@ -63,6 +125,32 @@ class TestClampFunction(unittest.TestCase):
|
||||
self.assertTrue(math.isnan(result))
|
||||
|
||||
|
||||
class TestMedianFunction(unittest.TestCase):
|
||||
|
||||
def test_odd_length_returns_middle_element(self) -> None:
|
||||
self.assertEqual(median([1, 3, 5]), 3)
|
||||
|
||||
def test_even_length_returns_float_average_of_two_middle_values(self) -> None:
|
||||
result = median([1, 2, 3, 4])
|
||||
self.assertIsInstance(result, float)
|
||||
self.assertEqual(result, 2.5)
|
||||
|
||||
def test_single_element_returns_that_element(self) -> None:
|
||||
self.assertEqual(median([7]), 7)
|
||||
|
||||
def test_empty_list_raises_value_error(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
median([])
|
||||
|
||||
def test_unsorted_input_sorts_correctly(self) -> None:
|
||||
self.assertEqual(median([3, 1, 2]), 2)
|
||||
|
||||
def test_unsorted_even_length_returns_float_average(self) -> None:
|
||||
result = median([10, 30, 20, 40])
|
||||
self.assertIsInstance(result, float)
|
||||
self.assertEqual(result, 25.0)
|
||||
|
||||
|
||||
class TestVarianceFunction(unittest.TestCase):
|
||||
|
||||
def test_empty_list_raises_value_error(self) -> None:
|
||||
@@ -94,3 +182,45 @@ class TestVarianceFunction(unittest.TestCase):
|
||||
self.assertEqual(variance((1, 2, 3, 4, 5)), 2.0)
|
||||
|
||||
|
||||
class TestPercentageFunction(unittest.TestCase):
|
||||
|
||||
def test_half_returns_50(self) -> None:
|
||||
self.assertEqual(percentage(50, 100), 50.0)
|
||||
|
||||
def test_quarter_returns_25(self) -> None:
|
||||
self.assertEqual(percentage(25, 100), 25.0)
|
||||
|
||||
def test_zero_value_returns_zero(self) -> None:
|
||||
self.assertEqual(percentage(0, 100), 0.0)
|
||||
|
||||
def test_value_exceeds_total(self) -> None:
|
||||
self.assertEqual(percentage(150, 100), 150.0)
|
||||
|
||||
def test_negative_value(self) -> None:
|
||||
self.assertEqual(percentage(-50, 100), -50.0)
|
||||
|
||||
def test_negative_total(self) -> None:
|
||||
self.assertEqual(percentage(50, -100), -50.0)
|
||||
|
||||
def test_both_negative(self) -> None:
|
||||
self.assertEqual(percentage(-50, -100), 50.0)
|
||||
|
||||
def test_float_inputs(self) -> None:
|
||||
result = percentage(33.0, 100.0)
|
||||
self.assertAlmostEqual(result, 33.0)
|
||||
|
||||
def test_total_zero_raises_value_error(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
percentage(50, 0)
|
||||
|
||||
def test_total_zero_float_raises_value_error(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
percentage(50.0, 0.0)
|
||||
|
||||
def test_integer_inputs_return_float(self) -> None:
|
||||
result = percentage(1, 4)
|
||||
self.assertIsInstance(result, float)
|
||||
self.assertEqual(result, 25.0)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,700 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
import io
|
||||
import json
|
||||
import unittest
|
||||
import urllib.error
|
||||
from typing import Any
|
||||
from unittest import mock
|
||||
|
||||
from typosaurus_sandbox.research.client import RsearchClient, RsearchError
|
||||
from typosaurus_sandbox.research.envelopes import (
|
||||
ChatResponse,
|
||||
ChatUsage,
|
||||
DeepReport,
|
||||
DescribeResponse,
|
||||
SearchGrade,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
|
||||
WEB_RESPONSE: dict[str, Any] = {
|
||||
"query": "asyncio python",
|
||||
"source": "duckduckgo",
|
||||
"count": 3,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"timestamp": "2026-08-07T12:00:00Z",
|
||||
"results": [
|
||||
{
|
||||
"title": "asyncio documentation",
|
||||
"url": "https://docs.python.org/3/library/asyncio.html",
|
||||
"description": "Asynchronous I/O event loop.",
|
||||
"source": "docs.python.org",
|
||||
"extra": {"rank": 1},
|
||||
"index": 0,
|
||||
},
|
||||
{
|
||||
"title": "asyncio in Python",
|
||||
"url": "https://example.com/asyncio",
|
||||
"description": "Tutorial on asyncio.",
|
||||
"source": "example.com",
|
||||
"extra": {"rank": 2},
|
||||
"index": 1,
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
AI_MEMORY_RESPONSE: dict[str, Any] = {
|
||||
"query": "python history",
|
||||
"source": "ai",
|
||||
"count": 0,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [],
|
||||
"ai_response": "From memory: Python was released in 1991 by Guido van Rossum.",
|
||||
"ai_error": None,
|
||||
}
|
||||
|
||||
AI_PROVIDER_RESPONSE: dict[str, Any] = {
|
||||
"query": "quantum computing",
|
||||
"source": "google",
|
||||
"count": 1,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [
|
||||
{
|
||||
"title": "Quantum computing overview",
|
||||
"url": "https://example.com/quantum",
|
||||
"description": "Overview of quantum computing.",
|
||||
"source": "example.com",
|
||||
"extra": {},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"ai_response": "Quantum computing uses qubits. [citation:1]",
|
||||
"ai_error": None,
|
||||
}
|
||||
|
||||
GRADED_RESPONSE: dict[str, Any] = {
|
||||
"query": "deep research",
|
||||
"source": "google",
|
||||
"count": 1,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [
|
||||
{
|
||||
"title": "Deep research systems",
|
||||
"url": "https://example.com/deep-research",
|
||||
"description": "Survey of deep research systems.",
|
||||
"source": "example.com",
|
||||
"extra": {},
|
||||
"index": 0,
|
||||
"grade": {
|
||||
"overall": 9.2,
|
||||
"relevance": 8.8,
|
||||
"depth": 9.0,
|
||||
"authority": 9.5,
|
||||
"freshness": 7.0,
|
||||
"word_count": 1200,
|
||||
"intent_hits": 4,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
DEEP_RESPONSE: dict[str, Any] = {
|
||||
"query": "deep research systems",
|
||||
"source": "google",
|
||||
"count": 8,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [
|
||||
{
|
||||
"title": "Deep research systems",
|
||||
"url": "https://example.com/deep-research",
|
||||
"description": "Survey of deep research systems.",
|
||||
"source": "example.com",
|
||||
"extra": {},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"deep": {
|
||||
"query": "deep research systems",
|
||||
"markdown": "# Deep research\n\nA survey.",
|
||||
"sources": [
|
||||
{
|
||||
"title": "Deep research systems",
|
||||
"url": "https://example.com/deep-research",
|
||||
"description": "Survey of deep research systems.",
|
||||
"source": "example.com",
|
||||
"extra": {},
|
||||
"grade": {
|
||||
"overall": 9.2,
|
||||
"relevance": 8.8,
|
||||
"depth": 9.0,
|
||||
"authority": 9.5,
|
||||
"freshness": 7.0,
|
||||
"word_count": 1200,
|
||||
"intent_hits": 4,
|
||||
},
|
||||
}
|
||||
],
|
||||
"graded_count": 8,
|
||||
"total_count": 10,
|
||||
"model": "gemma-3-12b-it",
|
||||
"elapsed": 166.96,
|
||||
"cache_hit": False,
|
||||
"rounds": 3,
|
||||
"queries_tried": ["deep research systems", "deep research architecture"],
|
||||
"error": None,
|
||||
},
|
||||
}
|
||||
|
||||
IMAGES_RESPONSE: dict[str, Any] = {
|
||||
"query": "aurora borealis",
|
||||
"source": "wikimedia",
|
||||
"count": 2,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [
|
||||
{
|
||||
"title": "Aurora borealis over Norway",
|
||||
"url": "https://commons.wikimedia.org/wiki/File:Aurora.jpg",
|
||||
"description": "Photograph of the aurora borealis.",
|
||||
"source": "wikimedia",
|
||||
"extra": {
|
||||
"thumbnail": "https://upload.wikimedia.org/thumb.jpg",
|
||||
"dimensions": {"width": 1920, "height": 1080},
|
||||
"mime": "image/jpeg",
|
||||
"license": "CC BY-SA 4.0",
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
CHAT_RESPONSE: dict[str, Any] = {
|
||||
"response": "The answer.",
|
||||
"prompt": "question",
|
||||
"json_mode": True,
|
||||
"cached": False,
|
||||
"error": None,
|
||||
"usage": {
|
||||
"prompt_tokens": 120,
|
||||
"completion_tokens": 80,
|
||||
"total_tokens": 200,
|
||||
"cost_usd": 0.0012,
|
||||
},
|
||||
}
|
||||
|
||||
DESCRIBE_RESPONSE: dict[str, Any] = {
|
||||
"url": "https://example.com/page",
|
||||
"description": "Page description",
|
||||
"elapsed": 1.23,
|
||||
"timestamp": "2026-08-07T12:00:00Z",
|
||||
}
|
||||
|
||||
SEARCH_EMPTY_OK: dict[str, Any] = {
|
||||
"query": "q",
|
||||
"source": "s",
|
||||
"count": 1,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [],
|
||||
}
|
||||
|
||||
|
||||
def _recorded_request(fixture: dict[str, Any]) -> tuple[list[tuple[Any, ...]], Any]:
|
||||
recorded: list[tuple[Any, ...]] = []
|
||||
|
||||
def fake(
|
||||
method: str,
|
||||
path: str,
|
||||
params: dict[str, str] | None,
|
||||
payload: bytes | None,
|
||||
headers: dict[str, str] | None,
|
||||
timeout: float | None,
|
||||
) -> tuple[int, dict[str, Any]]:
|
||||
recorded.append((method, path, params, payload, headers, timeout))
|
||||
return 200, fixture
|
||||
|
||||
return recorded, fake
|
||||
|
||||
|
||||
def _raising_request(message: str, status_code: int) -> Any:
|
||||
def fake(
|
||||
method: str,
|
||||
path: str,
|
||||
params: dict[str, str] | None,
|
||||
payload: bytes | None,
|
||||
headers: dict[str, str] | None,
|
||||
timeout: float | None,
|
||||
) -> tuple[int, dict[str, Any]]:
|
||||
raise RsearchError(message, status_code)
|
||||
|
||||
return fake
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, status: int, body: bytes) -> None:
|
||||
self.status = status
|
||||
self._body = body
|
||||
|
||||
def __enter__(self) -> "_FakeResponse":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def read(self) -> bytes:
|
||||
return self._body
|
||||
|
||||
|
||||
class TestSearchResponseParsing(unittest.TestCase):
|
||||
|
||||
def test_web_results_parse_into_search_response(self) -> None:
|
||||
response = SearchResponse.from_dict(WEB_RESPONSE)
|
||||
self.assertEqual(response.query, "asyncio python")
|
||||
self.assertEqual(response.source, "duckduckgo")
|
||||
self.assertEqual(response.count, 3)
|
||||
self.assertTrue(response.success)
|
||||
self.assertIsNone(response.error)
|
||||
self.assertEqual(response.timestamp, "2026-08-07T12:00:00Z")
|
||||
self.assertEqual(len(response.results), 2)
|
||||
first = response.results[0]
|
||||
self.assertIsInstance(first, SearchResult)
|
||||
self.assertEqual(first.title, "asyncio documentation")
|
||||
self.assertEqual(first.url, "https://docs.python.org/3/library/asyncio.html")
|
||||
self.assertEqual(first.description, "Asynchronous I/O event loop.")
|
||||
self.assertEqual(first.source, "docs.python.org")
|
||||
self.assertEqual(first.extra, {"rank": 1})
|
||||
self.assertEqual(first.index, 0)
|
||||
self.assertIsNone(first.content)
|
||||
self.assertIsNone(first.grade)
|
||||
self.assertIsNone(first.query_origin)
|
||||
self.assertIsNone(response.ai_response)
|
||||
self.assertIsNone(response.deep)
|
||||
|
||||
def test_ai_memory_variant_parses(self) -> None:
|
||||
response = SearchResponse.from_dict(AI_MEMORY_RESPONSE)
|
||||
self.assertEqual(response.source, "ai")
|
||||
self.assertEqual(response.results, [])
|
||||
self.assertIn("From memory", response.ai_response)
|
||||
self.assertIsNone(response.ai_error)
|
||||
|
||||
def test_ai_provider_variant_parses(self) -> None:
|
||||
response = SearchResponse.from_dict(AI_PROVIDER_RESPONSE)
|
||||
self.assertEqual(response.source, "google")
|
||||
self.assertEqual(len(response.results), 1)
|
||||
self.assertIn("[citation:1]", response.ai_response)
|
||||
self.assertIsNone(response.ai_error)
|
||||
|
||||
def test_deep_block_parses_into_deep_report(self) -> None:
|
||||
response = SearchResponse.from_dict(DEEP_RESPONSE)
|
||||
self.assertIsNotNone(response.deep)
|
||||
deep = response.deep
|
||||
self.assertIsInstance(deep, DeepReport)
|
||||
self.assertEqual(deep.query, "deep research systems")
|
||||
self.assertEqual(deep.markdown, "# Deep research\n\nA survey.")
|
||||
self.assertEqual(deep.graded_count, 8)
|
||||
self.assertEqual(deep.total_count, 10)
|
||||
self.assertEqual(deep.model, "gemma-3-12b-it")
|
||||
self.assertEqual(deep.elapsed, 166.96)
|
||||
self.assertFalse(deep.cache_hit)
|
||||
self.assertEqual(deep.rounds, 3)
|
||||
self.assertEqual(deep.queries_tried, ["deep research systems", "deep research architecture"])
|
||||
self.assertIsNone(deep.error)
|
||||
self.assertEqual(len(deep.sources), 1)
|
||||
source = deep.sources[0]
|
||||
self.assertIsInstance(source, SearchResult)
|
||||
self.assertEqual(source.url, "https://example.com/deep-research")
|
||||
self.assertIsInstance(source.grade, SearchGrade)
|
||||
self.assertEqual(source.grade.overall, 9.2)
|
||||
|
||||
def test_images_results_parse_extra_metadata(self) -> None:
|
||||
response = SearchResponse.from_dict(IMAGES_RESPONSE)
|
||||
self.assertEqual(response.source, "wikimedia")
|
||||
result = response.results[0]
|
||||
self.assertEqual(result.extra["mime"], "image/jpeg")
|
||||
self.assertEqual(result.extra["dimensions"], {"width": 1920, "height": 1080})
|
||||
self.assertEqual(result.extra["license"], "CC BY-SA 4.0")
|
||||
self.assertIn("thumbnail", result.extra)
|
||||
|
||||
def test_result_grade_parses_into_search_grade(self) -> None:
|
||||
response = SearchResponse.from_dict(GRADED_RESPONSE)
|
||||
grade = response.results[0].grade
|
||||
self.assertIsInstance(grade, SearchGrade)
|
||||
self.assertEqual(grade.overall, 9.2)
|
||||
self.assertEqual(grade.relevance, 8.8)
|
||||
self.assertEqual(grade.depth, 9.0)
|
||||
self.assertEqual(grade.authority, 9.5)
|
||||
self.assertEqual(grade.freshness, 7.0)
|
||||
self.assertEqual(grade.word_count, 1200)
|
||||
self.assertEqual(grade.intent_hits, 4)
|
||||
|
||||
def test_sparse_body_parses_with_defaults(self) -> None:
|
||||
response = SearchResponse.from_dict({"query": "x", "success": True})
|
||||
self.assertEqual(response.source, "")
|
||||
self.assertEqual(response.count, 0)
|
||||
self.assertEqual(response.results, [])
|
||||
self.assertIsNone(response.error)
|
||||
self.assertIsNone(response.deep)
|
||||
self.assertIsNone(response.ai_response)
|
||||
|
||||
|
||||
class TestSearchRequestConstruction(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
async def test_search_forwards_all_parameters(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(SEARCH_EMPTY_OK)
|
||||
client._request = fake
|
||||
await client.search(
|
||||
"query text",
|
||||
source="google",
|
||||
count=7,
|
||||
content=True,
|
||||
type="images",
|
||||
deep=True,
|
||||
ai=True,
|
||||
cache=False,
|
||||
)
|
||||
method, path, params, payload, headers, timeout = recorded[0]
|
||||
self.assertEqual(method, "GET")
|
||||
self.assertEqual(path, "/search")
|
||||
self.assertEqual(
|
||||
params,
|
||||
{
|
||||
"query": "query text",
|
||||
"source": "google",
|
||||
"count": "7",
|
||||
"content": "true",
|
||||
"type": "images",
|
||||
"deep": "true",
|
||||
"ai": "true",
|
||||
"cache": "false",
|
||||
},
|
||||
)
|
||||
self.assertIsNone(payload)
|
||||
self.assertIsNone(headers)
|
||||
self.assertEqual(timeout, 180.0)
|
||||
|
||||
async def test_search_without_deep_uses_request_timeout(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(SEARCH_EMPTY_OK)
|
||||
client._request = fake
|
||||
await client.search("q")
|
||||
self.assertEqual(recorded[0][5], 30.0)
|
||||
|
||||
async def test_count_none_omits_count_parameter(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(SEARCH_EMPTY_OK)
|
||||
client._request = fake
|
||||
await client.search("q")
|
||||
self.assertNotIn("count", recorded[0][2])
|
||||
|
||||
async def test_count_zero_forwarded_and_server_clamp_parsed(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(
|
||||
{"query": "q", "source": "s", "count": 1, "success": True, "error": None, "results": []}
|
||||
)
|
||||
client._request = fake
|
||||
response = await client.search("q", count=0)
|
||||
self.assertEqual(recorded[0][2]["count"], "0")
|
||||
self.assertEqual(response.count, 1)
|
||||
|
||||
async def test_count_above_limit_forwarded_and_server_clamp_parsed(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(
|
||||
{"query": "q", "source": "s", "count": 10, "success": True, "error": None, "results": []}
|
||||
)
|
||||
client._request = fake
|
||||
response = await client.search("q", count=25)
|
||||
self.assertEqual(recorded[0][2]["count"], "25")
|
||||
self.assertEqual(response.count, 10)
|
||||
|
||||
async def test_invalid_count_forwarded_and_server_clamp_parsed(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(
|
||||
{"query": "q", "source": "s", "count": 10, "success": True, "error": None, "results": []}
|
||||
)
|
||||
client._request = fake
|
||||
response = await client.search("q", count="not-a-number")
|
||||
self.assertEqual(recorded[0][2]["count"], "not-a-number")
|
||||
self.assertEqual(response.count, 10)
|
||||
|
||||
async def test_search_cache_hit_skips_second_request(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(SEARCH_EMPTY_OK)
|
||||
client._request = fake
|
||||
await client.search("cached query")
|
||||
await client.search("cached query")
|
||||
self.assertEqual(len(recorded), 1)
|
||||
|
||||
async def test_search_cache_disabled_repeats_request(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(SEARCH_EMPTY_OK)
|
||||
client._request = fake
|
||||
await client.search("uncached query", cache=False)
|
||||
await client.search("uncached query", cache=False)
|
||||
self.assertEqual(len(recorded), 2)
|
||||
|
||||
async def test_search_with_content_populates_content_cache(self) -> None:
|
||||
client = RsearchClient()
|
||||
fixture = {
|
||||
"query": "q",
|
||||
"source": "s",
|
||||
"count": 1,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"results": [
|
||||
{
|
||||
"title": "t",
|
||||
"url": "https://example.com/a",
|
||||
"description": "d",
|
||||
"source": "s",
|
||||
"extra": {},
|
||||
"content": "full page text",
|
||||
}
|
||||
],
|
||||
}
|
||||
recorded, fake = _recorded_request(fixture)
|
||||
client._request = fake
|
||||
await client.search("q", content=True)
|
||||
self.assertEqual(len(recorded), 1)
|
||||
self.assertEqual(client.get_cached_content("https://example.com/a"), "full page text")
|
||||
|
||||
async def test_search_error_in_body_surfaces_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
client._request = _raising_request("Empty query", 400)
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
await client.search("")
|
||||
self.assertEqual(ctx.exception.status_code, 400)
|
||||
self.assertEqual(str(ctx.exception), "Empty query")
|
||||
|
||||
|
||||
class TestChatResponseParsing(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
async def test_chat_response_parses_envelope(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(CHAT_RESPONSE)
|
||||
client._request = fake
|
||||
response = await client.chat("question", json_mode=True)
|
||||
method, path, params, payload, headers, timeout = recorded[0]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertEqual(path, "/chat")
|
||||
self.assertEqual(json.loads(payload), {"prompt": "question", "json": True})
|
||||
self.assertEqual(headers, {"Content-Type": "application/json"})
|
||||
self.assertIsNone(params)
|
||||
self.assertIsNone(timeout)
|
||||
self.assertIsInstance(response, ChatResponse)
|
||||
self.assertEqual(response.response, "The answer.")
|
||||
self.assertEqual(response.prompt, "question")
|
||||
self.assertTrue(response.json_mode)
|
||||
self.assertFalse(response.cached)
|
||||
self.assertIsNone(response.error)
|
||||
self.assertIsInstance(response.usage, ChatUsage)
|
||||
self.assertEqual(response.usage.prompt_tokens, 120)
|
||||
self.assertEqual(response.usage.completion_tokens, 80)
|
||||
self.assertEqual(response.usage.total_tokens, 200)
|
||||
self.assertEqual(response.usage.cost_usd, 0.0012)
|
||||
|
||||
async def test_chat_request_accepts_system_and_disables_cache(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(CHAT_RESPONSE)
|
||||
client._request = fake
|
||||
await client.chat("q", system="sys", cache=False)
|
||||
body = json.loads(recorded[0][3])
|
||||
self.assertEqual(body, {"prompt": "q", "system": "sys", "cache": False})
|
||||
|
||||
async def test_chat_error_raises_mapped_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
client._request = _raising_request("No prompt provided", 400)
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
await client.chat("")
|
||||
self.assertEqual(ctx.exception.status_code, 400)
|
||||
self.assertEqual(str(ctx.exception), "No prompt provided")
|
||||
|
||||
|
||||
class TestDescribeResponseParsing(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
async def test_describe_get_parses_envelope(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(DESCRIBE_RESPONSE)
|
||||
client._request = fake
|
||||
response = await client.describe("https://example.com/page")
|
||||
method, path, params, payload, headers, timeout = recorded[0]
|
||||
self.assertEqual(method, "GET")
|
||||
self.assertEqual(path, "/describe")
|
||||
self.assertEqual(params, {"url": "https://example.com/page"})
|
||||
self.assertIsNone(payload)
|
||||
self.assertIsNone(headers)
|
||||
self.assertIsNone(timeout)
|
||||
self.assertIsInstance(response, DescribeResponse)
|
||||
self.assertEqual(response.description, "Page description")
|
||||
self.assertEqual(response.url, "https://example.com/page")
|
||||
self.assertEqual(response.elapsed, 1.23)
|
||||
self.assertEqual(response.timestamp, "2026-08-07T12:00:00Z")
|
||||
self.assertIsNone(response.mime_type)
|
||||
self.assertIsNone(response.size)
|
||||
self.assertTrue(response.success)
|
||||
|
||||
async def test_describe_get_uses_cache(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(DESCRIBE_RESPONSE)
|
||||
client._request = fake
|
||||
await client.describe("https://example.com/page")
|
||||
await client.describe("https://example.com/page")
|
||||
self.assertEqual(len(recorded), 1)
|
||||
|
||||
async def test_describe_error_raises_mapped_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
client._request = _raising_request("No url provided", 400)
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
await client.describe("")
|
||||
self.assertEqual(ctx.exception.status_code, 400)
|
||||
self.assertEqual(str(ctx.exception), "No url provided")
|
||||
|
||||
async def test_describe_raw_posts_bytes_with_content_type(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(DESCRIBE_RESPONSE)
|
||||
client._request = fake
|
||||
image = b"\x89PNG\r\n\x1a\npayload"
|
||||
await client.describe_raw(image, mime_type="image/png")
|
||||
method, path, params, payload, headers, timeout = recorded[0]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertEqual(path, "/describe")
|
||||
self.assertEqual(payload, image)
|
||||
self.assertEqual(headers, {"Content-Type": "image/png"})
|
||||
|
||||
async def test_describe_upload_builds_multipart_body(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(DESCRIBE_RESPONSE)
|
||||
client._request = fake
|
||||
image = b"\x89PNGpayload"
|
||||
await client.describe_upload(image, filename="photo.png", mime_type="image/png")
|
||||
method, path, params, payload, headers, timeout = recorded[0]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertEqual(path, "/describe")
|
||||
self.assertIn(b'name="file"; filename="photo.png"', payload)
|
||||
self.assertIn(b"Content-Type: image/png", payload)
|
||||
self.assertIn(image, payload)
|
||||
self.assertIn("multipart/form-data; boundary=", headers["Content-Type"])
|
||||
|
||||
async def test_describe_raw_reuses_cache_by_hash(self) -> None:
|
||||
client = RsearchClient()
|
||||
recorded, fake = _recorded_request(DESCRIBE_RESPONSE)
|
||||
client._request = fake
|
||||
image = b"\x89PNGpayload"
|
||||
await client.describe_raw(image, mime_type="image/png")
|
||||
await client.describe_raw(image, mime_type="image/png")
|
||||
self.assertEqual(len(recorded), 1)
|
||||
|
||||
|
||||
class TestErrorInBodyHandling(unittest.TestCase):
|
||||
|
||||
def test_empty_query_error_in_body_maps_to_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
body = b'{"success": false, "error": "Empty query"}'
|
||||
error = urllib.error.HTTPError(
|
||||
"https://rsearch.app.molodetz.nl/search", 400, "Bad Request", {}, io.BytesIO(body)
|
||||
)
|
||||
with mock.patch("urllib.request.urlopen", side_effect=error):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": ""}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 400)
|
||||
self.assertEqual(str(ctx.exception), "Empty query")
|
||||
|
||||
def test_providers_exhausted_503_maps_to_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
body = b'{"success": false, "error": "All providers are exhausted, please try again later"}'
|
||||
error = urllib.error.HTTPError(
|
||||
"https://rsearch.app.molodetz.nl/search", 503, "Service Unavailable", {}, io.BytesIO(body)
|
||||
)
|
||||
with mock.patch("urllib.request.urlopen", side_effect=error):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 503)
|
||||
self.assertEqual(str(ctx.exception), "All providers are exhausted, please try again later")
|
||||
|
||||
def test_success_false_body_with_http_200_raises(self) -> None:
|
||||
client = RsearchClient()
|
||||
fake = _FakeResponse(200, b'{"success": false, "error": "Empty query"}')
|
||||
with mock.patch("urllib.request.urlopen", return_value=fake):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": ""}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 200)
|
||||
self.assertEqual(str(ctx.exception), "Empty query")
|
||||
|
||||
def test_detail_field_falls_back_for_error_message(self) -> None:
|
||||
client = RsearchClient()
|
||||
body = b'{"detail": "No url provided"}'
|
||||
error = urllib.error.HTTPError(
|
||||
"https://rsearch.app.molodetz.nl/describe", 400, "Bad Request", {}, io.BytesIO(body)
|
||||
)
|
||||
with mock.patch("urllib.request.urlopen", side_effect=error):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/describe", {"url": "x"}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 400)
|
||||
self.assertEqual(str(ctx.exception), "No url provided")
|
||||
|
||||
def test_title_field_falls_back_for_error_message(self) -> None:
|
||||
client = RsearchClient()
|
||||
body = b'{"title": "Provider error"}'
|
||||
error = urllib.error.HTTPError(
|
||||
"https://rsearch.app.molodetz.nl/search", 502, "Bad Gateway", {}, io.BytesIO(body)
|
||||
)
|
||||
with mock.patch("urllib.request.urlopen", side_effect=error):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 502)
|
||||
self.assertEqual(str(ctx.exception), "Provider error")
|
||||
|
||||
def test_empty_body_raises_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
fake = _FakeResponse(200, b"")
|
||||
with mock.patch("urllib.request.urlopen", return_value=fake):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 200)
|
||||
self.assertIn("empty response", str(ctx.exception))
|
||||
|
||||
def test_invalid_json_body_raises_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
fake = _FakeResponse(200, b"<html>not json</html>")
|
||||
with mock.patch("urllib.request.urlopen", return_value=fake):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 200)
|
||||
self.assertIn("invalid JSON", str(ctx.exception))
|
||||
|
||||
def test_non_dict_body_raises_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
fake = _FakeResponse(200, b'["not", "a", "dict"]')
|
||||
with mock.patch("urllib.request.urlopen", return_value=fake):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertEqual(ctx.exception.status_code, 200)
|
||||
|
||||
def test_connection_failure_raises_rsearch_error(self) -> None:
|
||||
client = RsearchClient()
|
||||
with mock.patch("urllib.request.urlopen", side_effect=urllib.error.URLError("connection refused")):
|
||||
with self.assertRaises(RsearchError) as ctx:
|
||||
client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertIn("connection failure", str(ctx.exception))
|
||||
|
||||
def test_successful_request_returns_status_and_body(self) -> None:
|
||||
client = RsearchClient()
|
||||
fake = _FakeResponse(200, b'{"success": true, "query": "q", "count": 1, "results": []}')
|
||||
with mock.patch("urllib.request.urlopen", return_value=fake) as urlopen:
|
||||
status, data = client._request("GET", "/search", {"query": "q"}, None, None, 30.0)
|
||||
self.assertEqual(status, 200)
|
||||
self.assertEqual(data, {"success": True, "query": "q", "count": 1, "results": []})
|
||||
request = urlopen.call_args[0][0]
|
||||
self.assertEqual(request.get_method(), "GET")
|
||||
self.assertEqual(request.get_full_url(), "https://rsearch.app.molodetz.nl/search?query=q")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user