Persist regenerated answer versions
This commit is contained in:
@@ -14,7 +14,7 @@ import httpx
|
||||
from fastapi import FastAPI, File, Form, Header, HTTPException, UploadFile
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from app import database
|
||||
from app import workspace
|
||||
@@ -69,6 +69,21 @@ class WorkspaceAgentRequest(AgentRequest):
|
||||
class StoredMessage(BaseModel):
|
||||
role: str = Field(pattern="^(user|assistant)$")
|
||||
content: str
|
||||
versions: list[str] = Field(default_factory=list, max_length=32)
|
||||
selected_version: int = Field(default=0, ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_answer_versions(self) -> "StoredMessage":
|
||||
if self.role == "user" and self.versions:
|
||||
raise ValueError("User messages cannot contain assistant answer versions.")
|
||||
if self.versions:
|
||||
if self.selected_version >= len(self.versions):
|
||||
raise ValueError("selected_version is outside the versions list.")
|
||||
if self.versions[self.selected_version] != self.content:
|
||||
raise ValueError("content must match the selected answer version.")
|
||||
elif self.selected_version != 0:
|
||||
raise ValueError("selected_version must be zero when versions are omitted.")
|
||||
return self
|
||||
|
||||
|
||||
class ConversationWrite(BaseModel):
|
||||
|
||||
Reference in New Issue
Block a user