fixed the errors- renamed jwt to authentication, used raw jwt, and removed search route

This commit is contained in:
sshah229 2026-02-07 03:14:30 -07:00
commit 0082627801
9 changed files with 27 additions and 105 deletions

View file

@ -1,7 +1,7 @@
"""
Authentication module for JWT-based auth with SQLite storage.
"""
from .jwt import create_access_token, get_current_subject, reload_secret
from .authentication import create_access_token, get_current_subject, reload_secret
from .storage import (
is_initialized,
create_initial_user,

View file

@ -4,7 +4,7 @@ from typing import Optional
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from jose import JWTError, jwt
import jwt
from .storage import load_jwt_secret
@ -72,7 +72,7 @@ async def get_current_subject(
detail="Invalid token payload",
)
return subject
except JWTError:
except jwt.InvalidTokenError:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid or expired token",

View file

@ -3,10 +3,10 @@ Unified core module for Unsloth backend
"""
# Inference
from .inference.inference import InferenceBackend, get_inference_backend
from .inference import InferenceBackend, get_inference_backend
# Training
from .training.training import UnslothTrainer, get_trainer, TrainingBackend, get_training_backend, create_training_handlers, TrainingProgress
from .training import UnslothTrainer, get_trainer, TrainingBackend, get_training_backend, create_training_handlers, TrainingProgress
# Configuration (from utils)
from utils.models import is_vision_model, ModelConfig, scan_trained_loras, load_model_defaults, get_base_model_from_lora

View file

@ -11,6 +11,7 @@ from .models import (
ModelDetails,
LoRAInfo,
LoRAScanResponse,
ModelListResponse,
)
from .auth import (
AuthSetupRequest,
@ -28,6 +29,7 @@ __all__ = [
"ModelDetails",
"LoRAInfo",
"LoRAScanResponse",
"ModelListResponse",
# Auth schemas
"AuthSetupRequest",
"AuthLoginRequest",

View file

@ -6,9 +6,11 @@ from typing import Optional, List, Dict, Any
class ModelDetails(BaseModel):
"""Detailed model configuration and metadata"""
model_name: str = Field(..., description="Model identifier")
config: Dict[str, Any] = Field(..., description="Model configuration dictionary")
"""Detailed model configuration and metadata - can be used for both list and detail views"""
id: str = Field(..., description="Model identifier")
model_name: Optional[str] = Field(None, description="Model identifier (alias for id, for backward compatibility)")
name: Optional[str] = Field(None, description="Display name for the model")
config: Optional[Dict[str, Any]] = Field(None, description="Model configuration dictionary")
is_vision: bool = Field(False, description="Whether model is a vision model")
is_lora: bool = Field(False, description="Whether model is a LoRA adapter")
base_model: Optional[str] = Field(None, description="Base model if this is a LoRA adapter")
@ -26,3 +28,9 @@ class LoRAScanResponse(BaseModel):
loras: List[LoRAInfo] = Field(default_factory=list, description="List of found LoRA adapters")
outputs_dir: str = Field(..., description="Directory that was scanned")
class ModelListResponse(BaseModel):
"""Response schema for listing models"""
models: List[ModelDetails] = Field(default_factory=list, description="List of models")
default_models: List[str] = Field(default_factory=list, description="List of default model IDs")

View file

@ -11,7 +11,7 @@ from models.auth import (
)
from models.users import Token
from auth import storage, hashing
from auth.jwt import create_access_token, reload_secret
from auth.authentication import create_access_token, reload_secret
router = APIRouter()
@ -56,7 +56,7 @@ async def setup_auth(payload: AuthSetupRequest) -> Token:
detail=f"Failed to create user: {str(e)}",
)
# Reload JWT secret from DB (so jwt.py picks it up)
# Reload JWT secret from DB (so authentication.py picks it up)
reload_secret()
# Issue a token for the new user

View file

@ -14,11 +14,10 @@ backend_path = Path(__file__).parent.parent.parent
if str(backend_path) not in sys.path:
sys.path.insert(0, str(backend_path))
from auth.jwt import get_current_subject
from auth.authentication import get_current_subject
# Import backend functions
try:
from utils.utils import search_hf_models
from utils.models import (
scan_trained_loras,
load_model_defaults,
@ -32,7 +31,6 @@ except ImportError:
parent_backend = backend_path.parent / "backend"
if str(parent_backend) not in sys.path:
sys.path.insert(0, str(parent_backend))
from utils.utils import search_hf_models
from utils.models import (
scan_trained_loras,
load_model_defaults,
@ -46,38 +44,8 @@ from models import (
ModelDetails,
LoRAScanResponse,
LoRAInfo,
ModelListResponse,
)
class ModelInfo(BaseModel):
"""Basic model info used in search/list responses"""
id: str
name: Optional[str] = None
is_vision: Optional[bool] = False
is_lora: Optional[bool] = False
class ModelSearchRequest(BaseModel):
"""Request body for model search"""
query: str
hf_token: Optional[str] = None
class ModelSearchResponse(BaseModel):
"""Response schema for model search"""
models: List[ModelInfo]
total: int
class ModelListResponse(BaseModel):
"""Response schema for listing models"""
models: List[ModelInfo]
default_models: List[str]
router = APIRouter()
logger = logging.getLogger(__name__)
@ -91,62 +59,6 @@ if not logger.handlers:
logger.setLevel(logging.INFO)
@router.post("/search")
async def search_models(
request: ModelSearchRequest,
current_subject: str = Depends(get_current_subject),
):
"""
Search for models on HuggingFace Hub.
This endpoint wraps the backend search_hf_models function.
"""
try:
# Call backend search function
gradio_update = search_hf_models(
search_query=request.query,
hf_token=request.hf_token
)
# Convert Gradio update to list of model IDs
model_list = []
if gradio_update and hasattr(gradio_update, 'choices'):
choices = gradio_update.choices
elif isinstance(gradio_update, dict) and 'choices' in gradio_update:
choices = gradio_update['choices']
elif isinstance(gradio_update, list):
choices = gradio_update
else:
choices = []
# Process choices - they may be tuples (display_name, model_id) or just strings
for choice in choices:
if isinstance(choice, tuple) and len(choice) >= 2:
# Format: (display_name, model_id)
model_id = choice[1] if len(choice) > 1 else choice[0]
display_name = choice[0]
model_info = ModelInfo(
id=model_id,
name=display_name
)
elif isinstance(choice, str):
# Just a model ID string
model_info = ModelInfo(id=choice)
else:
continue
model_list.append(model_info)
return ModelSearchResponse(
models=model_list,
total=len(model_list)
)
except Exception as e:
logger.error(f"Error searching models: {e}", exc_info=True)
raise HTTPException(
status_code=500,
detail=f"Failed to search models: {str(e)}"
)
@router.get("/list")
@ -167,7 +79,7 @@ async def list_models(
# Get loaded models
loaded_models = []
for model_name, model_data in inference_backend.models.items():
model_info = ModelInfo(
model_info = ModelDetails(
id=model_name,
name=model_name.split("/")[-1] if "/" in model_name else model_name,
is_vision=model_data.get("is_vision", False),
@ -182,7 +94,7 @@ async def list_models(
# Add default models
for model_id in default_models:
if model_id not in seen_ids:
model_info = ModelInfo(
model_info = ModelDetails(
id=model_id,
name=model_id.split("/")[-1] if "/" in model_id else model_id
)
@ -239,6 +151,7 @@ async def get_model_config(
pass
return ModelDetails(
id=model_name,
model_name=model_name,
config=config_dict,
is_vision=is_vision,

View file

@ -28,7 +28,7 @@ except ImportError:
from core.training import get_training_backend
# Auth
from auth.jwt import get_current_subject
from auth.authentication import get_current_subject
from models import (
TrainingStartRequest,

View file

@ -37,8 +37,7 @@ from .chat_templates import (
)
from .vlm_processing import generate_smart_vlm_instruction
from .data_collators import DeepSeekOCRDataCollator, VLMDataCollator
from .model_mappings import TEMPLATE_TO_MODEL_MAPPER
# , RESPONSE_MARKERS
from .model_mappings import TEMPLATE_TO_MODEL_MAPPER, RESPONSE_MARKERS
def check_dataset_format(dataset, is_vlm: bool = False) -> dict: