Skip to content
1 change: 1 addition & 0 deletions app/domain/rewriting_pipeline_execution_dto.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,4 +4,5 @@

class RewritingPipelineExecutionDTO(BaseModel):
execution: PipelineExecutionDTO
course_id: int = Field(alias="courseId")
to_be_rewritten: str = Field(alias="toBeRewritten")
5 changes: 5 additions & 0 deletions app/domain/status/rewriting_status_update_dto.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,10 @@
from typing import List

from app.domain.status.status_update_dto import StatusUpdateDTO


class RewritingStatusUpdateDTO(StatusUpdateDTO):
result: str = ""
suggestions: List[str] = []
inconsistencies: List[str] = []
improvement: str = ""
48 changes: 48 additions & 0 deletions app/pipeline/prompts/faq_consistency_prompt.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
faq_consistency_prompt = """
You are an AI assistant responsible for verifying the consistency of information.
### Task:
You have been provided with a list of FAQs and a final result. Your task is to determine whether the
final result is consistent with the given FAQs. Please compare each FAQ with the final result separately.

### Given FAQs:
{faqs}

### Final Result:
{final_result}

### Output:
Generate the following response dictionary:
"type": "consistent" or "inconsistent"

Firstly, identify the language of the course. The language of the course is either german or english. You can extract
the language from the existing FAQs. Your output should be in the same language as the course language.

The following four entries to the dictionary are optional and can only be set if inconsistencies are detected:
"faqs": This entry should be a list of Strings, each string represents an FAQ.
-Make sure each faq is separated by comma.
-Also end each faq with a newline character.
-The fields are exactly named faq_id, faq_question_title and faq_question_answer
and reside within properties dict of each list entry.
-Make sure to only include inconsistent faqs
-Do not include any additional FAQs that are consistent with the final_result.

"message": "The provided text was rephrased, however it contains inconsistent information with existing FAQs."
-Make sure to always insert two new lines after the last character of this sentences.
The affected FAQs can only contain the faq_id, faq_question_title, and faq_question_answer of inconsistent FAQs.
Make sure to not include any additional FAQs, that are consistent with the final_result.
Insert the faq_id, faq_question_title, and faq_question_answer of the inconsistent FAQ in the placeholder.

-"suggestion": This entry is a list of strings, each string represents a suggestion to improve the final result.\n
- Each suggestion should focus on a different inconsistency.
- Each suggestions highlights what is the inconsistency and how it can be improved.
- Do not mention the term final result, call it provided text
- Please ensure that at no time, you have a different amount of suggestions than inconsistencies.\n
Both should have the same amount of entries.

-"improved version": This entry should be a string that represents the improved version of the final result.

For each of the fields, make sure that the output is in the same language as the course language.
Comment thread
cremertim marked this conversation as resolved.
Outdated


Do NOT provide any explanations or additional text.
"""
32 changes: 0 additions & 32 deletions app/pipeline/prompts/faq_rewriting.py

This file was deleted.

92 changes: 86 additions & 6 deletions app/pipeline/rewriting_pipeline.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import json
import logging
from typing import Literal, Optional
from typing import Literal, Optional, List, Dict

from langchain.output_parsers import PydanticOutputParser
from langchain_core.prompts import (
Expand All @@ -12,10 +13,13 @@
from app.domain.rewriting_pipeline_execution_dto import RewritingPipelineExecutionDTO
from app.llm import CapabilityRequestHandler, RequirementList, CompletionArguments
from app.pipeline import Pipeline
from app.pipeline.prompts.faq_consistency_prompt import faq_consistency_prompt
from app.pipeline.prompts.rewriting_prompts import (
system_prompt_faq,
system_prompt_problem_statement,
)
from app.retrieval.faq_retrieval import FaqRetrieval
from app.vector_database.database import VectorDatabase
from app.web.status.status_update import RewritingCallback

logger = logging.getLogger(__name__)
Expand All @@ -38,8 +42,11 @@ def __init__(
context_length=16385,
)
)

self.db = VectorDatabase()
self.tokens = []
self.variant = variant
self.faq_retriever = FaqRetrieval(self.db.client)

def __call__(
self,
Expand All @@ -54,10 +61,10 @@ def __call__(
"faq": system_prompt_faq,
"problem_statement": system_prompt_problem_statement,
}
print(variant_prompts[self.variant])
prompt = variant_prompts[self.variant].format(
rewritten_text=dto.to_be_rewritten,
)

format_args = {"rewritten_text": dto.to_be_rewritten}

prompt = variant_prompts[self.variant].format(**format_args)
prompt = PyrisMessage(
sender=IrisMessageRole.SYSTEM,
contents=[TextMessageContentDTO(text_content=prompt)],
Expand All @@ -77,4 +84,77 @@ def __call__(
response = response.strip()

final_result = response
self.callback.done(final_result=final_result, tokens=self.tokens)
inconsistencies = []
improvement = ""
suggestions = []

if self.variant == "faq":
faqs = self.faq_retriever.get_faqs_from_db(
course_id=dto.course_id, search_text=response, result_limit=10
)
consistency_result = self.check_faq_consistency(faqs, final_result)

if "inconsistent" in consistency_result["type"].lower():
logging.warning("Detected inconsistencies in FAQ retrieval.")
inconsistencies = parse_inconsistencies(consistency_result["faqs"])
improvement = consistency_result["improved version"]
suggestions = consistency_result["suggestion"]

self.callback.done(
final_result=final_result,
tokens=self.tokens,
inconsistencies=inconsistencies,
improvement=improvement,
suggestions=suggestions,
)
Comment thread
cremertim marked this conversation as resolved.

def check_faq_consistency(
self, faqs: List[dict], final_result: str
) -> Dict[str, str]:
"""
Checks the consistency of the given FAQs with the provided final_result.
Returns "consistent" if there are no inconsistencies, otherwise returns "inconsistent".

:param faqs: List of retrieved FAQs.
:param final_result: The result to compare the FAQs against.

"""
properties_list = [entry["properties"] for entry in faqs]
Comment thread
cremertim marked this conversation as resolved.

consistency_prompt = faq_consistency_prompt.format(
faqs=properties_list, final_result=final_result
)
Comment thread
cremertim marked this conversation as resolved.

prompt = PyrisMessage(
sender=IrisMessageRole.SYSTEM,
contents=[TextMessageContentDTO(text_content=consistency_prompt)],
)

response = self.request_handler.chat(
[prompt], CompletionArguments(temperature=0.0), tools=None
)

self._append_tokens(response.token_usage, PipelineEnum.IRIS_REWRITING_PIPELINE)
result = response.contents[0].text_content
data = json.loads(result)

result_dict = {}

keys_to_check = ["type", "message", "faqs", "suggestion", "improved version"]

for key in keys_to_check:
if key in data:
result_dict[key] = data[key]

logging.info(f"Consistency FAQ consistency check response: {result_dict}")

return result_dict

Comment thread
cremertim marked this conversation as resolved.

def parse_inconsistencies(inconsistencies: List[Dict[str, str]]) -> List[str]:
logging.info("parse consistency")
parsed_inconsistencies = [
f"FAQ ID: {entry['faq_id']}, Title: {entry['faq_question_title']}, Answer: {entry['faq_question_answer']}"
for entry in inconsistencies
]
return parsed_inconsistencies
Comment thread
cremertim marked this conversation as resolved.
46 changes: 45 additions & 1 deletion app/retrieval/faq_retrieval.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@
from typing import List
from langsmith import traceable
from weaviate import WeaviateClient
from weaviate.collections.classes.filters import Filter

from app.common.PipelineEnum import PipelineEnum
from .basic_retrieval import BaseRetrieval, merge_retrieved_chunks
from ..common.pyris_message import PyrisMessage
Expand Down Expand Up @@ -44,7 +46,6 @@ def __call__(
base_url: str = None,
) -> List[dict]:
course_language = self.fetch_course_language(course_id)

response, response_hyde = self.run_parallel_rewrite_tasks(
chat_history=chat_history,
student_query=student_query,
Expand All @@ -67,3 +68,46 @@ def __call__(
for obj in response_hyde.objects
]
return merge_retrieved_chunks(basic_retrieved_faqs, hyde_retrieved_faqs)

def get_faqs_from_db(
self,
course_id: int,
search_text: str = None,
result_limit: int = 10,
hybrid_factor: float = 0.75,
) -> List[dict]:
"""
Retrieves FAQs directly from the database, optionally with a similarity search on question_title and question_answer.
Comment thread
cremertim marked this conversation as resolved.

:param course_id: ID of the course to fetch FAQs for a specific course.
:param search_text: Optional search text used for semantic search.
:param result_limit: Number of FAQs to return.
:param hybrid_factor: Weighting between vector-based and keyword-based results.
:return: List of retrieved FAQs.
"""
filter_weaviate = Filter.by_property("course_id").equal(course_id)

if search_text:
vec = self.llm_embedding.embed(search_text)

response = self.collection.query.hybrid(
query=search_text,
vector=vec,
alpha=hybrid_factor,
return_properties=self.get_schema_properties(),
limit=result_limit,
filters=filter_weaviate,
)
else:

response = self.collection.query.fetch_objects(
filters=filter_weaviate,
limit=result_limit,
return_properties=self.get_schema_properties(),
)

faqs = [
{"id": obj.uuid.int, "properties": obj.properties}
for obj in response.objects
]
return faqs
27 changes: 19 additions & 8 deletions app/web/status/status_update.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,8 @@ def done(
tokens: Optional[List[TokenUsageDTO]] = None,
next_stage_message: Optional[str] = None,
start_next_stage: bool = True,
inconsistencies: Optional[List[str]] = None,
improvement: Optional[str] = None,
):
"""
Transition the current stage to DONE and update the status.
Expand All @@ -122,6 +124,11 @@ def done(
self.status.tokens = tokens or self.status.tokens
if hasattr(self.status, "suggestions"):
self.status.suggestions = suggestions

if hasattr(self.status, "inconsistencies"):
self.status.inconsistencies = inconsistencies
if hasattr(self.status, "improvement"):
self.status.improvement = improvement
next_stage = self.get_next_stage()
if next_stage is not None:
self.stage = next_stage
Expand All @@ -133,6 +140,8 @@ def done(
self.status.result = None
if hasattr(self.status, "suggestions"):
self.status.suggestions = None
if hasattr(self.status, "inconsistencies"):
self.status.inconsistencies = None

def error(
self, message: str, exception=None, tokens: Optional[List[TokenUsageDTO]] = None
Expand Down Expand Up @@ -190,7 +199,7 @@ class CourseChatStatusCallback(StatusCallback):
def __init__(
self, run_id: str, base_url: str, initial_stages: List[StageDTO] = None
):
url = f"{base_url}/api/public/pyris/pipelines/course-chat/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/course-chat/runs/{run_id}/status"
current_stage_index = len(initial_stages) if initial_stages else 0
stages = initial_stages or []
stages += [
Expand All @@ -212,7 +221,7 @@ class ExerciseChatStatusCallback(StatusCallback):
def __init__(
self, run_id: str, base_url: str, initial_stages: List[StageDTO] = None
):
url = f"{base_url}/api/public/pyris/pipelines/tutor-chat/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/tutor-chat/runs/{run_id}/status"
current_stage_index = len(initial_stages) if initial_stages else 0
stages = initial_stages or []
stages += [
Expand All @@ -234,7 +243,7 @@ class ChatGPTWrapperStatusCallback(StatusCallback):
def __init__(
self, run_id: str, base_url: str, initial_stages: List[StageDTO] = None
):
url = f"{base_url}/api/public/pyris/pipelines/tutor-chat/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/tutor-chat/runs/{run_id}/status"
current_stage_index = len(initial_stages) if initial_stages else 0
stages = initial_stages or []
stages += [
Expand All @@ -256,7 +265,7 @@ def __init__(
base_url: str,
initial_stages: List[StageDTO],
):
url = f"{base_url}/api/public/pyris/pipelines/text-exercise-chat/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/text-exercise-chat/runs/{run_id}/status"
stages = initial_stages or []
stage = len(stages)
stages += [
Expand Down Expand Up @@ -287,7 +296,7 @@ def __init__(
base_url: str,
initial_stages: List[StageDTO],
):
url = f"{base_url}/api/public/pyris/pipelines/competency-extraction/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/competency-extraction/runs/{run_id}/status"
stages = initial_stages or []
stages.append(
StageDTO(
Expand All @@ -308,7 +317,9 @@ def __init__(
base_url: str,
initial_stages: List[StageDTO],
):
url = f"{base_url}/api/public/pyris/pipelines/rewriting/runs/{run_id}/status"
url = (
f"{base_url}/api/iris/public/pyris/pipelines/rewriting/runs/{run_id}/status"
)
stages = initial_stages or []
stages.append(
StageDTO(
Expand All @@ -329,7 +340,7 @@ def __init__(
base_url: str,
initial_stages: List[StageDTO],
):
url = f"{base_url}/api/public/pyris/pipelines/inconsistency-check/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/inconsistency-check/runs/{run_id}/status"
stages = initial_stages or []
stages.append(
StageDTO(
Expand All @@ -350,7 +361,7 @@ def __init__(
base_url: str,
initial_stages: List[StageDTO],
):
url = f"{base_url}/api/public/pyris/pipelines/lecture-chat/runs/{run_id}/status"
url = f"{base_url}/api/iris/public/pyris/pipelines/lecture-chat/runs/{run_id}/status"
stages = initial_stages or []
stage = len(stages)
stages += [
Expand Down