Files
enlace/tests/test_distill.py
T
msaldain 4048936067 Corregir lo que reportó ruff, que nunca se había ejecutado
ruff estaba configurado en pyproject.toml desde el primer commit y jamás se
había corrido. Tenía 15 hallazgos. Dos importan más allá del estilo:

- zip() sin strict= trunca en silencio al más corto. En las comparaciones de
  lotes eso significa que un test podía pasar sin haber comparado todo. Donde
  los largos deben coincidir ahora es strict=True; donde difieren a propósito
  (pares consecutivos) queda strict=False, que documenta la intención.
- Un import sin usar delataba algo peor: BraveBackend se había escrito sin una
  sola prueba. Se agregan ocho, contra una respuesta con la forma que devuelve
  la API, incluidas la limpieza de etiquetas, el caso de límite de tasa —que
  tiene que distinguirse de 'no respondió'— y que la credencial viaje en la
  cabecera y nunca en la URL. El respaldo tiene que funcionar justo cuando el
  primario ya falló; merecía la misma cobertura.

El resto es orden de imports, collections.abc y líneas largas.
2026-07-28 07:25:51 -03:00

329 lines
11 KiB
Python

"""Generación por lotes: idempotencia, reanudación y emparejado por custom_id.
Todo corre contra un cliente falso. Lo que hay que verificar es el harness —
que no duplique costo, que no pierda resultados y que no desalinee el dataset —
no el SDK de Anthropic.
"""
from __future__ import annotations
import json
from collections.abc import Iterator
from pathlib import Path
from typing import Any
import pytest
from enlace.config.load import load_distill_config
from enlace.config.schema import DistillConfig
from enlace.data.distill import (
DistillError,
Peticion,
Trabajo,
collect,
leer_especificaciones,
resumen_de_uso,
status,
submit,
)
class ClienteFalso:
"""Cliente de lotes en memoria, con fallos y demoras programables."""
def __init__(
self,
*,
fallan: set[str] | None = None,
rondas_hasta_terminar: int = 0,
desordenar: bool = True,
) -> None:
self.lotes: dict[str, list[dict[str, Any]]] = {}
self.fallan = fallan or set()
self.rondas_hasta_terminar = rondas_hasta_terminar
self.desordenar = desordenar
self.consultas = 0
self.envios = 0
def construir_params(self, peticion: Peticion) -> dict[str, Any]:
return {"model": "falso", "user": peticion.user}
def crear(self, pedidos: list[dict[str, Any]]) -> str:
self.envios += 1
batch_id = f"batch-{self.envios:03d}"
self.lotes[batch_id] = pedidos
return batch_id
def estado(self, batch_id: str) -> dict[str, Any]:
self.consultas += 1
listo = self.consultas > self.rondas_hasta_terminar * len(self.lotes)
return {
"processing_status": "ended" if listo else "in_progress",
"succeeded": len(self.lotes[batch_id]),
"errored": 0,
"processing": 0,
"canceled": 0,
"expired": 0,
}
def resultados(self, batch_id: str) -> Iterator[dict[str, Any]]:
pedidos = self.lotes[batch_id]
# La API devuelve los resultados en cualquier orden: el doble los
# invierte a propósito para que un emparejado por posición falle.
if self.desordenar:
pedidos = list(reversed(pedidos))
for pedido in pedidos:
cid = pedido["custom_id"]
if cid in self.fallan:
yield {"custom_id": cid, "type": "errored", "error_type": "api_error"}
else:
yield {
"custom_id": cid,
"type": "succeeded",
"text": f"respuesta para {pedido['params']['user']}",
"stop_reason": "end_turn",
"usage": {
"input_tokens": 100,
"output_tokens": 200,
"cache_read_input_tokens": 900,
"cache_creation_input_tokens": 0,
},
}
def cancelar(self, batch_id: str) -> None:
self.lotes.pop(batch_id, None)
@pytest.fixture
def config() -> DistillConfig:
return DistillConfig(requests_per_batch=3, poll_interval_seconds=5.0)
@pytest.fixture
def trabajo(tmp_path: Path) -> Trabajo:
return Trabajo(nombre="prueba", raiz=tmp_path)
def _peticiones(n: int) -> list[Peticion]:
return [
Peticion(system="Sos ENLACE.", user=f"consulta {i}", meta={"semilla": i}) for i in range(n)
]
def _sin_espera(trabajo, cliente, config, **kw):
return collect(trabajo, cliente, config, ahora=lambda: 0.0, dormir=lambda _: None, **kw)
# --- custom_id -------------------------------------------------------------
def test_el_custom_id_se_deriva_del_contenido():
"""Idempotencia: la misma especificación produce el mismo identificador,
así que re-generar un trabajo no vuelve a pagar lo ya resuelto."""
a = Peticion(system="S", user="U")
b = Peticion(system="S", user="U", meta={"otra": "cosa"})
assert a.custom_id == b.custom_id # los metadatos no van a la API
c = Peticion(system="S", user="OTRA")
assert a.custom_id != c.custom_id
def test_las_peticiones_duplicadas_se_descartan(trabajo):
guardadas = trabajo.guardar_peticiones([*_peticiones(3), *_peticiones(3)])
assert len(guardadas) == 3
# --- submit ----------------------------------------------------------------
def test_submit_parte_en_lotes_del_tamano_configurado(trabajo, config):
cliente = ClienteFalso()
ids = submit(trabajo, _peticiones(7), cliente, config)
assert len(ids) == 3 # 3 + 3 + 1
assert cliente.envios == 3
def test_submit_persiste_cada_lote_al_enviarlo(trabajo, config):
"""Si el proceso muere en el segundo envío, el primero no queda huérfano."""
class ClienteQueFalla(ClienteFalso):
def crear(self, pedidos):
if self.envios >= 1:
raise RuntimeError("caída simulada")
return super().crear(pedidos)
with pytest.raises(RuntimeError):
submit(trabajo, _peticiones(7), ClienteQueFalla(), config)
assert len(trabajo.leer_lotes()) == 1 # el primero quedó registrado
def test_submit_no_reenvia_lo_ya_enviado(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(6), cliente, config)
submit(trabajo, _peticiones(6), cliente, config)
assert cliente.envios == 2 # los mismos dos lotes, no cuatro
def test_submit_no_reenvia_lo_ya_resuelto(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(6), cliente, config)
_sin_espera(trabajo, cliente, config)
# Un trabajo ampliado: 6 viejas ya resueltas + 2 nuevas.
cliente2 = ClienteFalso()
submit(trabajo, _peticiones(8), cliente2, config)
enviados = sum(len(p) for p in cliente2.lotes.values())
assert enviados == 2
# --- collect ---------------------------------------------------------------
def test_collect_empareja_por_custom_id_no_por_posicion(trabajo, config):
"""El cliente falso devuelve los resultados invertidos a propósito.
Emparejar por índice mezclaría las respuestas en silencio, y un dataset mal
alineado es peor que uno vacío porque parece correcto.
"""
cliente = ClienteFalso(desordenar=True)
peticiones = _peticiones(3)
submit(trabajo, peticiones, cliente, config)
_sin_espera(trabajo, cliente, config)
por_id = {}
with trabajo.resultados_path.open(encoding="utf-8") as fh:
for linea in fh:
fila = json.loads(linea)
por_id[fila["custom_id"]] = fila["text"]
for p in peticiones:
assert por_id[p.custom_id] == f"respuesta para {p.user}"
def test_los_fallos_van_a_errors_y_no_se_pierden(trabajo, config):
peticiones = _peticiones(4)
cliente = ClienteFalso(fallan={peticiones[1].custom_id})
submit(trabajo, peticiones, cliente, config)
resumen = _sin_espera(trabajo, cliente, config)
assert resumen["ok"] == 3
assert resumen["error"] == 1
errores = trabajo.errores_path.read_text().strip().splitlines()
assert json.loads(errores[0])["custom_id"] == peticiones[1].custom_id
def test_recoger_dos_veces_no_duplica(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(4), cliente, config)
primero = _sin_espera(trabajo, cliente, config)
segundo = _sin_espera(trabajo, cliente, config)
assert primero["ok"] == 4
assert segundo["ok"] == 0
assert segundo["ya_estaban"] == 4
assert len(trabajo.resultados_path.read_text().strip().splitlines()) == 4
def test_collect_espera_a_los_lotes_en_curso(trabajo, config):
cliente = ClienteFalso(rondas_hasta_terminar=2)
submit(trabajo, _peticiones(3), cliente, config)
dormidas = []
resumen = collect(trabajo, cliente, config, ahora=lambda: 0.0, dormir=dormidas.append)
assert resumen["ok"] == 3
assert dormidas # efectivamente esperó
def test_collect_respeta_el_tope_de_espera(trabajo, config):
cliente = ClienteFalso(rondas_hasta_terminar=1000)
submit(trabajo, _peticiones(3), cliente, config)
reloj = iter([0.0, 1e9, 1e9])
with pytest.raises(DistillError, match="max_wait_hours"):
collect(trabajo, cliente, config, ahora=lambda: next(reloj), dormir=lambda _: None)
def test_collect_sin_lotes_da_un_error_util(trabajo, config):
with pytest.raises(DistillError, match="submit"):
_sin_espera(trabajo, ClienteFalso(), config)
def test_status_reporta_cada_lote(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(7), cliente, config)
estados = status(trabajo, cliente)
assert len(estados) == 3
assert all("processing_status" in e for e in estados)
# --- costo -----------------------------------------------------------------
def test_el_resumen_de_uso_aplica_el_descuento_y_la_cache(trabajo, config):
cliente = ClienteFalso()
submit(trabajo, _peticiones(2), cliente, config)
_sin_espera(trabajo, cliente, config)
uso = resumen_de_uso(trabajo, config)
assert uso["resultados"] == 2
assert uso["output_tokens"] == 400
assert uso["cache_read_input_tokens"] == 1800
# (200 entrada + 1800*0.1 caché)/1e6*5 + 400/1e6*25, todo al 50%.
esperado = ((200 + 180) / 1e6 * 5 + 400 / 1e6 * 25) * 0.5
assert uso["usd_estimado"] == pytest.approx(round(esperado, 4))
def test_sin_resultados_el_costo_es_cero(trabajo, config):
assert resumen_de_uso(trabajo, config)["usd_estimado"] == 0.0
# --- entrada ---------------------------------------------------------------
def test_lee_especificaciones_jsonl(tmp_path):
path = tmp_path / "peticiones.jsonl"
path.write_text(
json.dumps({"system": "S", "user": "U", "meta": {"k": 1}}, ensure_ascii=False)
+ "\n\n"
+ json.dumps({"system": "S", "user": "U2", "schema": {"type": "object"}})
+ "\n",
encoding="utf-8",
)
peticiones = leer_especificaciones(path)
assert len(peticiones) == 2 # la línea en blanco se ignora
assert peticiones[0].meta == {"k": 1}
assert peticiones[1].schema == {"type": "object"}
def test_especificacion_invalida_indica_la_linea(tmp_path):
path = tmp_path / "malas.jsonl"
path.write_text('{"system": "S", "user": "U"}\n{"system": "sin user"}\n')
with pytest.raises(DistillError, match=":2:"):
leer_especificaciones(path)
def test_manifiesto_faltante_da_un_error_util(trabajo):
with pytest.raises(DistillError, match="manifiesto"):
trabajo.leer_peticiones()
def test_ida_y_vuelta_del_manifiesto(trabajo):
originales = _peticiones(3)
trabajo.guardar_peticiones(originales)
assert trabajo.leer_peticiones() == originales
# --- config ----------------------------------------------------------------
def test_la_config_por_defecto_es_valida():
cfg = load_distill_config()
assert cfg.model == "claude-opus-5"
assert 0 < cfg.batch_discount <= 1
assert cfg.requests_per_batch <= 100_000
def test_rechaza_lotes_mayores_al_tope_de_la_api():
with pytest.raises(ValueError):
DistillConfig(requests_per_batch=200_000)