"""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 pathlib import Path from typing import Any, Iterator 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)