diff --git a/main.py b/main.py index 7c7b65b..fec87cd 100644 --- a/main.py +++ b/main.py @@ -14,6 +14,11 @@ app = FastAPI() s3 = boto3.client('s3') BUCKET = os.getenv('BUCKET_NAME', 'ia-gondola-projeto-2024') +# 'ambiente' vira componente de paths locais (/tmp/..., runs/detect/...) e de +# chaves S3 — sem allowlist, um valor tipo "../../../app" permitiria path +# traversal (leitura/escrita/exclusão fora do diretório esperado). +AMBIENTES_VALIDOS = {"gondola", "etiqueta"} + MODELOS_CARREGADOS = {} _treinamento_status = {"status": "idle", "ambiente": None, "versao": None, "detalhe": None} @@ -187,18 +192,19 @@ def _executar_treino_sync(ambiente: str, pular_triagem: bool, forcar_do_zero: bo modelo_base.add_callback("on_fit_epoch_end", _on_fit_epoch_end) # Fine-tuning: LR baixo preserva o que o modelo já aprendeu. # Primeiro treino: LR padrão + mais épocas para convergir do zero. - # project/name/exist_ok fixos: sem isso o Ultralytics incrementa o - # diretório a cada treino (train, train2, train3...) e o caminho abaixo - # ficava sempre apontando pro "train" original, subindo pro S3 um - # modelo de uma rodada antiga em vez do que acabou de treinar. - train_dir = f"runs/detect/{ambiente}" + # name/exist_ok fixos: sem isso o Ultralytics incrementa o diretório a + # cada treino (train, train2, train3...). Não passar "project" — o + # Ultralytics já resolve o project default a partir do runs_dir das + # settings, e forçar um valor aqui duplicava o caminho + # (runs/detect/runs/detect/{ambiente}), fazendo o `best` abaixo nunca + # ser encontrado. Lemos o caminho real direto do trainer, sem adivinhar. if is_fine_tuning: - modelo_base.train(data=yaml_path, epochs=20, imgsz=640, batch=8, device='cpu', plots=True, lr0=0.001, lrf=0.01, project="runs/detect", name=ambiente, exist_ok=True) + modelo_base.train(data=yaml_path, epochs=20, imgsz=640, batch=8, device='cpu', plots=True, lr0=0.001, lrf=0.01, name=ambiente, exist_ok=True) else: - modelo_base.train(data=yaml_path, epochs=30, imgsz=640, batch=8, device='cpu', plots=True, lr0=0.01, lrf=0.01, project="runs/detect", name=ambiente, exist_ok=True) + modelo_base.train(data=yaml_path, epochs=30, imgsz=640, batch=8, device='cpu', plots=True, lr0=0.01, lrf=0.01, name=ambiente, exist_ok=True) modelo_base.reset_callbacks() - best = f"{train_dir}/weights/best.pt" + best = str(modelo_base.trainer.save_dir / "weights" / "best.pt") if os.path.exists(best): carimbo = datetime.now().strftime("%Y%m%d_%H%M") s3.upload_file(best, BUCKET, f"modelos/{ambiente}/atual/cerebro.pt") @@ -246,6 +252,8 @@ def _treinar_bg(ambiente: str, pular_triagem: bool, forcar_do_zero: bool = False @app.post("/detectar") async def detectar(ambiente: str = Form(...), file: UploadFile = File(...)): + if ambiente not in AMBIENTES_VALIDOS: + raise HTTPException(status_code=400, detail=f"ambiente inválido: {ambiente}") try: Image.MAX_IMAGE_PIXELS = None modelo = carregar_modelo_do_s3(ambiente) @@ -265,10 +273,12 @@ async def detectar(ambiente: str = Form(...), file: UploadFile = File(...)): @app.post("/treinar") async def treinar(dados: dict): global _treinamento_status + ambiente = dados.get("ambiente", "gondola") + if ambiente not in AMBIENTES_VALIDOS: + raise HTTPException(status_code=400, detail=f"ambiente inválido: {ambiente}") with _status_lock: if _treinamento_status["status"] == "running": raise HTTPException(status_code=409, detail="Treinamento já em andamento") - ambiente = dados.get("ambiente", "gondola") pular_triagem = dados.get("pular_triagem", False) forcar_do_zero = dados.get("forcar_do_zero", False) _treinamento_status = {"status": "running", "ambiente": ambiente, "versao": None, "detalhe": None}