Files
spawner1145-comfyui-lsnet/backend_lsnet/api.py
T
2026-10-06 00:30:52 +08:00

200 lines
8.4 KiB
Python

import base64
import logging
from typing import Callable
from threading import Lock
from secrets import compare_digest
from io import BytesIO
import asyncio
import concurrent.futures
import os
import glob
import json
from fastapi import FastAPI, Depends, HTTPException
from fastapi.security import HTTPBasic, HTTPBasicCredentials
from pydantic import BaseModel, Field
from PIL import Image
import numpy as np
from backend_lsnet.inference import process_image_from_pil
from backend_lsnet.analysis_api import (FeaturesRequest, AnalysisRequest, FeatureToolsRequest,
extract_request, serialized_cache, analysis_request, tools_request)
try:
from modules import shared
from modules.call_queue import queue_lock as webui_queue_lock
IN_WEBUI = True
except ImportError:
IN_WEBUI = False
shared = type('Shared', (), {'cmd_opts': type('CmdOpts', (), {'api_auth': None})()})()
webui_queue_lock = None
# 设置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
from backend_lsnet.model_paths import (
get_available_checkpoints, get_available_csv, get_checkpoint_path, get_class_csv,
)
class InferenceRequest(BaseModel):
input_image: str = Field(..., description="Input image as Base64 encoded string")
model_name: str = Field('Kaloscope', description="Model name (subfolder in models/kaloscope/)")
device: str = Field('cuda', description="Device to use")
top_k: int = Field(5, ge=1, le=20, description="Number of top predictions")
threshold: float = Field(0.0, ge=0.0, le=1.0, description="Probability threshold")
mode: str = 'auto'
output_type: str = 'default'
layers: str = '-1'
intermediate_norm: bool = True
class InferenceResponse(BaseModel):
results: dict = Field(..., description="Inference results")
info: str = Field(..., description="Additional information")
class CancelResponse(BaseModel):
info: str = Field(..., description="Cancel operation result")
class Api:
def __init__(self, app: FastAPI, queue_lock: Lock = None, prefix: str = "/kaloscope/v1"):
self.app = app
self.queue_lock = queue_lock or Lock()
self.prefix = prefix
self.credentials = {}
self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=1)
if IN_WEBUI and shared.cmd_opts.api_auth:
for auth in shared.cmd_opts.api_auth.split(","):
user, password = auth.split(":")
self.credentials[user] = password
self.add_api_route(
"infer",
self.endpoint_infer,
methods=["POST"],
response_model=InferenceResponse,
summary="Perform artist style inference",
description="Classify an image or extract features with LSNet or DINOv3."
)
self.add_api_route(
"cancel",
self.endpoint_cancel,
methods=["POST"],
response_model=CancelResponse,
summary="Cancel the current inference task",
description="Terminates the ongoing inference task."
)
self.add_api_route('features', self.endpoint_features, methods=['POST'], summary='Extract a batch once; return reusable NPZ cache')
self.add_api_route('analyze', self.endpoint_analyze, methods=['POST'], summary='Render any analysis chart from features/cache or an image batch')
self.add_api_route('feature-tools', self.endpoint_feature_tools, methods=['POST'], summary='Common features, similarity or group comparison without inference')
self.add_api_route('models', self.endpoint_models, methods=['GET'], summary='Available model folders and supported chart/feature types')
def auth(self, creds: HTTPBasicCredentials = Depends(HTTPBasic())):
if not self.credentials:
return True
if creds.username in self.credentials:
if compare_digest(creds.password, self.credentials[creds.username]):
return True
raise HTTPException(
status_code=401,
detail="Incorrect username or password",
headers={"WWW-Authenticate": "Basic"}
)
def add_api_route(self, path: str, endpoint: Callable, **kwargs):
route = f"{self.prefix}/{path}" if self.prefix else path
dependencies = [Depends(self.auth)] if self.credentials else []
self.app.add_api_route(route, endpoint, dependencies=dependencies, **kwargs)
def decode_base64_image(self, base64_str: str) -> Image.Image:
try:
img_data = base64.b64decode(base64_str, validate=True)
img = Image.open(BytesIO(img_data))
img.load()
return img
except base64.binascii.Error:
raise HTTPException(400, "Invalid Base64 string format")
except Exception as e:
raise HTTPException(400, f"Failed to decode image: {str(e)}")
async def run_inference(self, image, **kwargs):
"""Run inference in a separate thread"""
loop = asyncio.get_event_loop()
try:
return await loop.run_in_executor(self.executor, lambda: process_image_from_pil(image, **kwargs))
except Exception as e:
logger.error(f"Inference execution failed: {str(e)}")
raise
async def endpoint_infer(self, req: InferenceRequest):
logger.info(f"Received inference request: model_name={req.model_name}")
try:
with self.queue_lock:
input_image = self.decode_base64_image(req.input_image)
checkpoints = get_available_checkpoints(req.model_name)
if not checkpoints:
raise HTTPException(400, f"No checkpoints found for model {req.model_name}")
checkpoint = get_checkpoint_path(req.model_name)
if not os.path.exists(checkpoint):
raise HTTPException(400, f"Checkpoint not found: {checkpoint}")
class_csv = get_class_csv(req.model_name)
infer_args = {
"checkpoint": checkpoint,
"mode": req.mode,
"output_type": req.output_type,
"layers": req.layers,
"intermediate_norm": req.intermediate_norm,
"device": req.device,
"top_k": req.top_k,
"threshold": req.threshold,
"class_csv": class_csv
}
# Run inference
results = await self.run_inference(input_image, **infer_args)
return InferenceResponse(results=results, info="Inference completed successfully")
except HTTPException:
raise
except Exception as e:
logger.error(f"Inference failed: {str(e)}")
raise HTTPException(500, f"Inference failed: {str(e)}")
async def endpoint_cancel(self):
# For simplicity, just return a message since inference is quick
return CancelResponse(info="No active inference to cancel")
async def run_feature_job(self, callback):
def locked_job():
with self.queue_lock:
return callback()
try:
return await asyncio.get_running_loop().run_in_executor(self.executor, locked_job)
except HTTPException:
raise
except (ValueError, FileNotFoundError, TypeError, KeyError) as error:
raise HTTPException(400, str(error)) from error
async def endpoint_features(self, req: FeaturesRequest):
return await self.run_feature_job(lambda: serialized_cache(extract_request(req, self.decode_base64_image)[0]))
async def endpoint_analyze(self, req: AnalysisRequest):
return await self.run_feature_job(lambda: analysis_request(req, self.decode_base64_image))
async def endpoint_feature_tools(self, req: FeatureToolsRequest):
return await self.run_feature_job(lambda: tools_request(req))
async def endpoint_models(self):
from backend_lsnet.model_paths import get_available_models
from feature_analysis import CHART_TYPES, TENSOR_LAYOUTS
from model_loading import FEATURE_OUTPUTS
return {'models': get_available_models(), 'feature_outputs': FEATURE_OUTPUTS,
'chart_types': CHART_TYPES, 'tensor_layouts': TENSOR_LAYOUTS}
def on_app_started(demo, app):
"""Called when the webui app starts"""
queue_lock = webui_queue_lock or Lock()
api = Api(app, queue_lock)
logger.info("Kaloscope API routes added to webui")