From 5ac985e435222d6ad01763fd40052733593319ac Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Andreas=20K=C3=B6pf?= Date: Fri, 16 Dec 2022 14:09:48 +0100 Subject: [PATCH] fix typo --- backend/app/api/v1/tasks2.py | 17 ++++++++++++++++- backend/app/models/payload_column_type.py | 2 +- backend/app/prompt_repository.py | 18 +++++++++--------- 3 files changed, 26 insertions(+), 11 deletions(-) diff --git a/backend/app/api/v1/tasks2.py b/backend/app/api/v1/tasks2.py index 77b30ef0..a334d0b3 100644 --- a/backend/app/api/v1/tasks2.py +++ b/backend/app/api/v1/tasks2.py @@ -4,7 +4,7 @@ from typing import Any from uuid import UUID from app.api import deps -from app.prompt_repository import PromptRepository, TaskPayload +from app.prompt_repository import PromptRepository, RateSummaryPayload, TaskPayload from app.schemas import protocol as protocol_schema from fastapi import APIRouter, Depends, HTTPException from fastapi.security.api_key import APIKey @@ -179,6 +179,21 @@ def post_interaction( f"Frontend reports rating of {interaction.post_id=} with {interaction.rating=} by {interaction.user=}." ) # check if rating in range + + work_package = pr.fetch_workpackage_by_postid(interaction.post_id) + work_payload: RateSummaryPayload = work_package.payload.payload + if ( + type(work_payload) != RateSummaryPayload + or interaction.rating < work_payload.scale.min + or interaction.rating > work_payload.scale.max + ): + raise HTTPException( + status_code=HTTP_400_BAD_REQUEST, + detail="Invalid response type.", + ) + + pr.store_rating(interaction) + # here we would store the rating in the database return protocol_schema.TaskDone( reply_to_post_id=interaction.post_id, diff --git a/backend/app/models/payload_column_type.py b/backend/app/models/payload_column_type.py index a95ccb53..fbda51ce 100644 --- a/backend/app/models/payload_column_type.py +++ b/backend/app/models/payload_column_type.py @@ -14,7 +14,7 @@ payload_type_registry = {} P = TypeVar("P", bound=BaseModel) -def payload_tpye(cls: Type[P]) -> Type[P]: +def payload_type(cls: Type[P]) -> Type[P]: payload_type_registry[cls.__name__] = cls return cls diff --git a/backend/app/prompt_repository.py b/backend/app/prompt_repository.py index 24a9ea1c..85f372d3 100644 --- a/backend/app/prompt_repository.py +++ b/backend/app/prompt_repository.py @@ -4,24 +4,24 @@ from typing import Literal, Optional from uuid import UUID, uuid4 from app.models import ApiClient, Person, Post, PostReaction, WorkPackage -from app.models.payload_column_type import PayloadContainer, payload_tpye +from app.models.payload_column_type import PayloadContainer, payload_type from app.schemas import protocol as protocol_schema from pydantic import BaseModel from sqlmodel import Session -@payload_tpye +@payload_type class TaskPayload(BaseModel): type: str -@payload_tpye +@payload_type class SummarizationStoryPayload(TaskPayload): type: Literal["summarize_story"] = "summarize_story" story: str -@payload_tpye +@payload_type class RateSummaryPayload(TaskPayload): type: Literal["rate_summary"] = "rate_summary" full_text: str @@ -29,31 +29,31 @@ class RateSummaryPayload(TaskPayload): scale: protocol_schema.RatingScale -@payload_tpye +@payload_type class InitialPromptPayload(TaskPayload): type: Literal["initial_prompt"] = "initial_prompt" hint: str -@payload_tpye +@payload_type class UserReplyPayload(TaskPayload): type: Literal["user_reply"] = "user_reply" conversation: protocol_schema.Conversation hint: str | None -@payload_tpye +@payload_type class AssistantReplyPayload(TaskPayload): type: Literal["assistant_reply"] = "assistant_reply" conversation: protocol_schema.Conversation -@payload_tpye +@payload_type class ReactionPayload(BaseModel): type: str -@payload_tpye +@payload_type class RatingReactionPayload(ReactionPayload): type: Literal["post_rating"] = "post_rating" rating: str