Add OasstError exception class and exception filter

This commit is contained in:
Andreas Köpf
2022-12-28 14:10:15 +01:00
parent 99bf737117
commit dda668bcd5
7 changed files with 105 additions and 42 deletions
+10
View File
@@ -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(
+7 -3
View File
@@ -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)
+15 -21
View File
@@ -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,
)
+2 -1
View File
@@ -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)
+26
View File
@@ -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
+21
View File
@@ -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})'
+24 -17
View File
@@ -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(