mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-09-09 11:15:08 +08:00
Add OasstError exception class and exception filter
This commit is contained in:
@@ -7,10 +7,20 @@ import fastapi
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
from oasst_backend.api.v1.api import api_router
|
from oasst_backend.api.v1.api import api_router
|
||||||
from oasst_backend.config import settings
|
from oasst_backend.config import settings
|
||||||
|
from oasst_backend.exceptions import OasstError
|
||||||
from starlette.middleware.cors import CORSMiddleware
|
from starlette.middleware.cors import CORSMiddleware
|
||||||
|
|
||||||
app = fastapi.FastAPI(title=settings.PROJECT_NAME, openapi_url=f"{settings.API_V1_STR}/openapi.json")
|
app = fastapi.FastAPI(title=settings.PROJECT_NAME, openapi_url=f"{settings.API_V1_STR}/openapi.json")
|
||||||
|
|
||||||
|
|
||||||
|
@app.exception_handler(OasstError)
|
||||||
|
async def http_exception_handler(request: fastapi.Request, ex: OasstError):
|
||||||
|
logger.error(f"{request.method} {request.url} failed: {repr(ex)}")
|
||||||
|
return fastapi.responses.JSONResponse(
|
||||||
|
status_code=ex.http_status_code, content={"message": ex.message, "error_code": ex.error_code}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# Set all CORS enabled origins
|
# Set all CORS enabled origins
|
||||||
if settings.BACKEND_CORS_ORIGINS:
|
if settings.BACKEND_CORS_ORIGINS:
|
||||||
app.add_middleware(
|
app.add_middleware(
|
||||||
|
|||||||
@@ -3,11 +3,12 @@ from secrets import token_hex
|
|||||||
from typing import Generator
|
from typing import Generator
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import HTTPException, Security
|
from fastapi import Security
|
||||||
from fastapi.security.api_key import APIKey, APIKeyHeader, APIKeyQuery
|
from fastapi.security.api_key import APIKey, APIKeyHeader, APIKeyQuery
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from oasst_backend.config import settings
|
from oasst_backend.config import settings
|
||||||
from oasst_backend.database import engine
|
from oasst_backend.database import engine
|
||||||
|
from oasst_backend.exceptions import OasstError, error_codes
|
||||||
from oasst_backend.models import ApiClient
|
from oasst_backend.models import ApiClient
|
||||||
from sqlmodel import Session
|
from sqlmodel import Session
|
||||||
from starlette.status import HTTP_403_FORBIDDEN
|
from starlette.status import HTTP_403_FORBIDDEN
|
||||||
@@ -36,9 +37,12 @@ def api_auth(
|
|||||||
api_key: APIKey,
|
api_key: APIKey,
|
||||||
db: Session,
|
db: Session,
|
||||||
) -> ApiClient:
|
) -> ApiClient:
|
||||||
|
|
||||||
if api_key is None and not settings.DEBUG_SKIP_API_KEY_CHECK:
|
if api_key is None and not settings.DEBUG_SKIP_API_KEY_CHECK:
|
||||||
raise HTTPException(status_code=HTTP_403_FORBIDDEN, detail="Could not validate credentials")
|
raise OasstError(
|
||||||
|
"Could not validate credentials",
|
||||||
|
error_code=error_codes.API_CLIENT_NOT_AUTHORIZED,
|
||||||
|
http_status_code=HTTP_403_FORBIDDEN,
|
||||||
|
)
|
||||||
|
|
||||||
if settings.DEBUG_SKIP_API_KEY_CHECK or settings.DEBUG_ALLOW_ANY_API_KEY:
|
if settings.DEBUG_SKIP_API_KEY_CHECK or settings.DEBUG_ALLOW_ANY_API_KEY:
|
||||||
# make sure that a dummy api key exits in db (foreign key references)
|
# make sure that a dummy api key exits in db (foreign key references)
|
||||||
|
|||||||
@@ -3,14 +3,14 @@ import random
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends
|
||||||
from fastapi.security.api_key import APIKey
|
from fastapi.security.api_key import APIKey
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from oasst_backend.api import deps
|
from oasst_backend.api import deps
|
||||||
|
from oasst_backend.exceptions import OasstError, error_codes
|
||||||
from oasst_backend.prompt_repository import PromptRepository
|
from oasst_backend.prompt_repository import PromptRepository
|
||||||
from oasst_shared.schemas import protocol as protocol_schema
|
from oasst_shared.schemas import protocol as protocol_schema
|
||||||
from sqlmodel import Session
|
from sqlmodel import Session
|
||||||
from starlette.status import HTTP_400_BAD_REQUEST
|
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
@@ -114,10 +114,7 @@ def generate_task(request: protocol_schema.TaskRequest) -> protocol_schema.Task:
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
case _:
|
case _:
|
||||||
raise HTTPException(
|
raise OasstError("Invalid request type", error_codes.TASK_INVALID_REQUEST_TYPE)
|
||||||
status_code=HTTP_400_BAD_REQUEST,
|
|
||||||
detail="Invalid request type.",
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.info(f"Generated {task=}.")
|
logger.info(f"Generated {task=}.")
|
||||||
|
|
||||||
@@ -134,6 +131,7 @@ def request_task(
|
|||||||
"""
|
"""
|
||||||
Create new task.
|
Create new task.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
api_client = deps.api_auth(api_key, db)
|
api_client = deps.api_auth(api_key, db)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -142,11 +140,11 @@ def request_task(
|
|||||||
pr = PromptRepository(db, api_client, request.user)
|
pr = PromptRepository(db, api_client, request.user)
|
||||||
pr.store_task(task)
|
pr.store_task(task)
|
||||||
|
|
||||||
|
except OasstError:
|
||||||
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to generate task.")
|
logger.exception("Failed to generate task..")
|
||||||
raise HTTPException(
|
raise OasstError("Failed to generate task.", error_codes.TASK_GENERATION_FAILED)
|
||||||
status_code=HTTP_400_BAD_REQUEST,
|
|
||||||
)
|
|
||||||
return task
|
return task
|
||||||
|
|
||||||
|
|
||||||
@@ -171,11 +169,11 @@ def acknowledge_task(
|
|||||||
logger.info(f"Frontend acknowledges task {task_id=}, {ack_request=}.")
|
logger.info(f"Frontend acknowledges task {task_id=}, {ack_request=}.")
|
||||||
pr.bind_frontend_post_id(task_id=task_id, post_id=ack_request.post_id)
|
pr.bind_frontend_post_id(task_id=task_id, post_id=ack_request.post_id)
|
||||||
|
|
||||||
|
except OasstError:
|
||||||
|
raise
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to acknowledge task.")
|
logger.exception("Failed to acknowledge task.")
|
||||||
raise HTTPException(
|
raise OasstError("Failed to acknowledge task.", error_codes.TASK_ACK_FAILED)
|
||||||
status_code=HTTP_400_BAD_REQUEST,
|
|
||||||
)
|
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|
||||||
@@ -242,13 +240,9 @@ def post_interaction(
|
|||||||
# here we would store the ranking in the database
|
# here we would store the ranking in the database
|
||||||
return protocol_schema.TaskDone()
|
return protocol_schema.TaskDone()
|
||||||
case _:
|
case _:
|
||||||
raise HTTPException(
|
raise OasstError("Invalid response type.", error_codes.TASK_INVALID_RESPONSE_TYPE)
|
||||||
status_code=HTTP_400_BAD_REQUEST,
|
except OasstError:
|
||||||
detail="Invalid response type.",
|
raise
|
||||||
)
|
|
||||||
|
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Interaction request failed.")
|
logger.exception("Interaction request failed.")
|
||||||
raise HTTPException(
|
raise OasstError("Interaction request failed.", error_codes.TASK_INTERACTION_REQUEST_FAILED)
|
||||||
status_code=HTTP_400_BAD_REQUEST,
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -1,8 +1,9 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
from oasst_backend.config import settings
|
from oasst_backend.config import settings
|
||||||
|
from oasst_backend.exceptions import OasstError, error_codes
|
||||||
from sqlmodel import create_engine
|
from sqlmodel import create_engine
|
||||||
|
|
||||||
if settings.DATABASE_URI is None:
|
if settings.DATABASE_URI is None:
|
||||||
raise ValueError("DATABASE_URI is not set")
|
raise OasstError("DATABASE_URI is not set", error_code=error_codes.DATABASE_URI_NOT_SET)
|
||||||
|
|
||||||
engine = create_engine(settings.DATABASE_URI)
|
engine = create_engine(settings.DATABASE_URI)
|
||||||
|
|||||||
@@ -0,0 +1,26 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""
|
||||||
|
Open-Assistant backend API error codes.
|
||||||
|
"""
|
||||||
|
# 0-1000: general errors
|
||||||
|
GENERIC_ERROR = 0
|
||||||
|
DATABASE_URI_NOT_SET = 1
|
||||||
|
API_CLIENT_NOT_AUTHORIZED = 2
|
||||||
|
|
||||||
|
# 1000-2000: tasks endpoint
|
||||||
|
TASK_INVALID_REQUEST_TYPE = 1000
|
||||||
|
TASK_ACK_FAILED = 1001
|
||||||
|
TASK_INVALID_RESPONSE_TYPE = 1002
|
||||||
|
TASK_INTERACTION_REQUEST_FAILED = 1003
|
||||||
|
TASK_GENERATION_FAILED = 1004
|
||||||
|
|
||||||
|
# 2000-3000: prompt_repository
|
||||||
|
INVALID_POST_ID = 2000
|
||||||
|
POST_NOT_FOUND = 2001
|
||||||
|
RATING_OUT_OF_RANGE = 2002
|
||||||
|
INVALID_RANKING_VALUE = 2003
|
||||||
|
WORK_PACKAGE_NOT_FOUND = 2004
|
||||||
|
WORK_PACKAGE_EXPIRED = 2005
|
||||||
|
WORK_PACKAGE_PAYLOAD_TYPE_MISMATCH = 2006
|
||||||
|
INVALID_TASK_TYPE = 2007
|
||||||
|
USER_NOT_SPECIFIED = 2008
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
import oasst_backend.error_codes as error_codes # noqa: F401
|
||||||
|
from starlette.status import HTTP_400_BAD_REQUEST
|
||||||
|
|
||||||
|
|
||||||
|
class OasstError(Exception):
|
||||||
|
"""Base class for Open-Assistant exceptions."""
|
||||||
|
|
||||||
|
message: str
|
||||||
|
error_code: int
|
||||||
|
http_status_code: int
|
||||||
|
|
||||||
|
def __init__(self, message: str, error_code: int, http_status_code: int = HTTP_400_BAD_REQUEST):
|
||||||
|
super().__init__(message, error_code, http_status_code) # make excetpion picklable (fill args member)
|
||||||
|
self.message = message
|
||||||
|
self.error_code = error_code
|
||||||
|
self.http_status_code = http_status_code
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
class_name = self.__class__.__name__
|
||||||
|
return f'{class_name}(message="{self.message}", error_code={self.error_code}, http_status_code={self.http_status_code})'
|
||||||
@@ -5,6 +5,7 @@ from uuid import UUID, uuid4
|
|||||||
|
|
||||||
import oasst_backend.models.db_payload as db_payload
|
import oasst_backend.models.db_payload as db_payload
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
from oasst_backend.exceptions import OasstError, error_codes
|
||||||
from oasst_backend.journal_writer import JournalWriter
|
from oasst_backend.journal_writer import JournalWriter
|
||||||
from oasst_backend.models import ApiClient, Person, Post, PostReaction, TextLabels, WorkPackage
|
from oasst_backend.models import ApiClient, Person, Post, PostReaction, TextLabels, WorkPackage
|
||||||
from oasst_backend.models.payload_column_type import PayloadContainer
|
from oasst_backend.models.payload_column_type import PayloadContainer
|
||||||
@@ -52,9 +53,9 @@ class PromptRepository:
|
|||||||
|
|
||||||
def validate_post_id(self, post_id: str) -> None:
|
def validate_post_id(self, post_id: str) -> None:
|
||||||
if not isinstance(post_id, str):
|
if not isinstance(post_id, str):
|
||||||
raise TypeError(f"post_id must be string, not {type(post_id)}")
|
raise OasstError(f"post_id must be string, not {type(post_id)}", error_codes.INVALID_POST_ID)
|
||||||
if not post_id:
|
if not post_id:
|
||||||
raise ValueError("post_id must not be empty")
|
raise OasstError("post_id must not be empty", error_codes.INVALID_POST_ID)
|
||||||
|
|
||||||
def bind_frontend_post_id(self, task_id: UUID, post_id: str):
|
def bind_frontend_post_id(self, task_id: UUID, post_id: str):
|
||||||
self.validate_post_id(post_id)
|
self.validate_post_id(post_id)
|
||||||
@@ -66,9 +67,9 @@ class PromptRepository:
|
|||||||
.first()
|
.first()
|
||||||
)
|
)
|
||||||
if work_pack is None:
|
if work_pack is None:
|
||||||
raise KeyError(f"WorkPackage for task {task_id} not found")
|
raise OasstError(f"WorkPackage for task {task_id} not found", error_codes.WORK_PACKAGE_NOT_FOUND)
|
||||||
if work_pack.expiry_date is not None and datetime.utcnow() > work_pack.expiry_date:
|
if work_pack.expiry_date is not None and datetime.utcnow() > work_pack.expiry_date:
|
||||||
raise RuntimeError("WorkPackage already expired.")
|
raise OasstError("WorkPackage already expired.", error_codes.WORK_PACKAGE_EXPIRED)
|
||||||
|
|
||||||
# ToDo: check race-condition, transaction
|
# ToDo: check race-condition, transaction
|
||||||
|
|
||||||
@@ -105,7 +106,7 @@ class PromptRepository:
|
|||||||
.one_or_none()
|
.one_or_none()
|
||||||
)
|
)
|
||||||
if fail_if_missing and post is None:
|
if fail_if_missing and post is None:
|
||||||
raise KeyError(f"Post with post_id {frontend_post_id} not found.")
|
raise OasstError(f"Post with post_id {frontend_post_id} not found.", error_codes.POST_NOT_FOUND)
|
||||||
return post
|
return post
|
||||||
|
|
||||||
def fetch_workpackage_by_postid(self, post_id: str) -> WorkPackage:
|
def fetch_workpackage_by_postid(self, post_id: str) -> WorkPackage:
|
||||||
@@ -134,7 +135,7 @@ class PromptRepository:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if parent_post is None:
|
if parent_post is None:
|
||||||
raise KeyError(f"Post for post_id {reply.post_id} not found.")
|
raise OasstError(f"Post for post_id {reply.post_id} not found.", error_codes.POST_NOT_FOUND)
|
||||||
|
|
||||||
# create reply post
|
# create reply post
|
||||||
user_post_id = uuid4()
|
user_post_id = uuid4()
|
||||||
@@ -156,12 +157,15 @@ class PromptRepository:
|
|||||||
work_package = self.fetch_workpackage_by_postid(rating.post_id)
|
work_package = self.fetch_workpackage_by_postid(rating.post_id)
|
||||||
work_payload: db_payload.RateSummaryPayload = work_package.payload.payload
|
work_payload: db_payload.RateSummaryPayload = work_package.payload.payload
|
||||||
if type(work_payload) != db_payload.RateSummaryPayload:
|
if type(work_payload) != db_payload.RateSummaryPayload:
|
||||||
raise ValueError(
|
raise OasstError(
|
||||||
f"work_package payload type mismatch: {type(work_payload)=} != {db_payload.RateSummaryPayload}"
|
f"work_package payload type mismatch: {type(work_payload)=} != {db_payload.RateSummaryPayload}",
|
||||||
|
error_codes.WORK_PACKAGE_PAYLOAD_TYPE_MISMATCH,
|
||||||
)
|
)
|
||||||
|
|
||||||
if rating.rating < work_payload.scale.min or rating.rating > work_payload.scale.max:
|
if rating.rating < work_payload.scale.min or rating.rating > work_payload.scale.max:
|
||||||
raise ValueError(f"Invalid rating value: {rating.rating=} not in {work_payload.scale=}")
|
raise OasstError(
|
||||||
|
f"Invalid rating value: {rating.rating=} not in {work_payload.scale=}", error_codes.RATING_OUT_OF_RANGE
|
||||||
|
)
|
||||||
|
|
||||||
# store reaction to post
|
# store reaction to post
|
||||||
reaction_payload = db_payload.RatingReactionPayload(rating=rating.rating)
|
reaction_payload = db_payload.RatingReactionPayload(rating=rating.rating)
|
||||||
@@ -185,8 +189,9 @@ class PromptRepository:
|
|||||||
# validate ranking
|
# validate ranking
|
||||||
num_replies = len(work_payload.replies)
|
num_replies = len(work_payload.replies)
|
||||||
if sorted(ranking.ranking) != list(range(num_replies)):
|
if sorted(ranking.ranking) != list(range(num_replies)):
|
||||||
raise ValueError(
|
raise OasstError(
|
||||||
f"Invalid ranking submitted. Each reply index must appear exactly once ({num_replies=})."
|
f"Invalid ranking submitted. Each reply index must appear exactly once ({num_replies=}).",
|
||||||
|
error_codes.INVALID_RANKING_VALUE,
|
||||||
)
|
)
|
||||||
|
|
||||||
# store reaction to post
|
# store reaction to post
|
||||||
@@ -201,8 +206,9 @@ class PromptRepository:
|
|||||||
case db_payload.RankInitialPromptsPayload:
|
case db_payload.RankInitialPromptsPayload:
|
||||||
# validate ranking
|
# validate ranking
|
||||||
if sorted(ranking.ranking) != list(range(num_prompts := len(work_payload.prompts))):
|
if sorted(ranking.ranking) != list(range(num_prompts := len(work_payload.prompts))):
|
||||||
raise ValueError(
|
raise OasstError(
|
||||||
f"Invalid ranking submitted. Each reply index must appear exactly once ({num_prompts=})."
|
f"Invalid ranking submitted. Each reply index must appear exactly once ({num_prompts=}).",
|
||||||
|
error_codes.INVALID_RANKING_VALUE,
|
||||||
)
|
)
|
||||||
|
|
||||||
# store reaction to post
|
# store reaction to post
|
||||||
@@ -215,8 +221,9 @@ class PromptRepository:
|
|||||||
return reaction
|
return reaction
|
||||||
|
|
||||||
case _:
|
case _:
|
||||||
raise ValueError(
|
raise OasstError(
|
||||||
f"work_package payload type mismatch: {type(work_payload)=} != {db_payload.RankConversationRepliesPayload}"
|
f"work_package payload type mismatch: {type(work_payload)=} != {db_payload.RankConversationRepliesPayload}",
|
||||||
|
error_codes.WORK_PACKAGE_PAYLOAD_TYPE_MISMATCH,
|
||||||
)
|
)
|
||||||
|
|
||||||
def store_task(self, task: protocol_schema.Task) -> WorkPackage:
|
def store_task(self, task: protocol_schema.Task) -> WorkPackage:
|
||||||
@@ -253,7 +260,7 @@ class PromptRepository:
|
|||||||
)
|
)
|
||||||
|
|
||||||
case _:
|
case _:
|
||||||
raise ValueError(f"Invalid task type: {type(task)=}")
|
raise OasstError(f"Invalid task type: {type(task)=}", error_codes.INVALID_TASK_TYPE)
|
||||||
|
|
||||||
wp = self.insert_work_package(payload=payload, id=task.id)
|
wp = self.insert_work_package(payload=payload, id=task.id)
|
||||||
assert wp.id == task.id
|
assert wp.id == task.id
|
||||||
@@ -310,7 +317,7 @@ class PromptRepository:
|
|||||||
|
|
||||||
def insert_reaction(self, post_id: UUID, payload: db_payload.ReactionPayload) -> PostReaction:
|
def insert_reaction(self, post_id: UUID, payload: db_payload.ReactionPayload) -> PostReaction:
|
||||||
if self.person_id is None:
|
if self.person_id is None:
|
||||||
raise ValueError("User required")
|
raise OasstError("User required", error_codes.USER_NOT_SPECIFIED)
|
||||||
|
|
||||||
container = PayloadContainer(payload=payload)
|
container = PayloadContainer(payload=payload)
|
||||||
reaction = PostReaction(
|
reaction = PostReaction(
|
||||||
|
|||||||
Reference in New Issue
Block a user