Samuel Oberhofer commited on
Commit ·
2322128
1
Parent(s): ff27905
feat: Implement retriever and guardrail modules
Browse files- guards/input.py +30 -0
- guards/output.py +27 -0
- rag/retriever.py +47 -0
guards/input.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import re
|
| 2 |
+
|
| 3 |
+
def is_valid(query: str) -> bool:
|
| 4 |
+
"""
|
| 5 |
+
Validates the user's query.
|
| 6 |
+
"""
|
| 7 |
+
# Check for query length
|
| 8 |
+
if len(query) > 500:
|
| 9 |
+
return False
|
| 10 |
+
|
| 11 |
+
# Check for SQL injection patterns
|
| 12 |
+
sql_injection_patterns = [
|
| 13 |
+
r"(\s*(--|#|;))",
|
| 14 |
+
r"(\s*(union|select|insert|update|delete|drop|alter)\s+)",
|
| 15 |
+
]
|
| 16 |
+
for pattern in sql_injection_patterns:
|
| 17 |
+
if re.search(pattern, query, re.IGNORECASE):
|
| 18 |
+
return False
|
| 19 |
+
|
| 20 |
+
return True
|
| 21 |
+
|
| 22 |
+
if __name__ == '__main__':
|
| 23 |
+
# Example usage
|
| 24 |
+
valid_query = "What are the computer science courses?"
|
| 25 |
+
invalid_query_long = "a" * 501
|
| 26 |
+
invalid_query_sql = "SELECT * FROM students;"
|
| 27 |
+
|
| 28 |
+
print(f"'{valid_query}' is valid: {is_valid(valid_query)}")
|
| 29 |
+
print(f"'long query' is valid: {is_valid(invalid_query_long)}")
|
| 30 |
+
print(f"'{invalid_query_sql}' is valid: {is_valid(invalid_query_sql)}")
|
guards/output.py
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import re
|
| 2 |
+
|
| 3 |
+
def sanitize(response: str) -> str:
|
| 4 |
+
"""
|
| 5 |
+
Sanitizes the LLM's response.
|
| 6 |
+
"""
|
| 7 |
+
# Redact email addresses
|
| 8 |
+
response = re.sub(r"[a-zA-Z0-9_.+-]+@[a-zA-Z0-9-]+\.[a-zA-Z0-9-.]+", "[REDACTED EMAIL]", response)
|
| 9 |
+
|
| 10 |
+
# Check for hallucination patterns
|
| 11 |
+
hallucination_patterns = [
|
| 12 |
+
r"as an ai language model",
|
| 13 |
+
r"i cannot answer that question",
|
| 14 |
+
]
|
| 15 |
+
for pattern in hallucination_patterns:
|
| 16 |
+
if re.search(pattern, response, re.IGNORECASE):
|
| 17 |
+
return "I'm sorry, I don't have enough information to answer that question."
|
| 18 |
+
|
| 19 |
+
return response
|
| 20 |
+
|
| 21 |
+
if __name__ == '__main__':
|
| 22 |
+
# Example usage
|
| 23 |
+
valid_response = "Professor Smith's email is john.smith@university.edu."
|
| 24 |
+
hallucination_response = "As an AI language model, I cannot provide that information."
|
| 25 |
+
|
| 26 |
+
print(f"Original: '{valid_response}'\nSanitized: '{sanitize(valid_response)}'")
|
| 27 |
+
print(f"\nOriginal: '{hallucination_response}'\nSanitized: '{sanitize(hallucination_response)}'")
|
rag/retriever.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import chromadb
|
| 2 |
+
from sentence_transformers import SentenceTransformer
|
| 3 |
+
import sqlite3
|
| 4 |
+
|
| 5 |
+
def search(query: str, top_k: int = 5):
|
| 6 |
+
"""
|
| 7 |
+
Searches the vector store for the most relevant documents to a given query.
|
| 8 |
+
"""
|
| 9 |
+
# Handle special case for listing all students
|
| 10 |
+
if "give me the names of the students" in query.lower():
|
| 11 |
+
conn = sqlite3.connect('database/university.db')
|
| 12 |
+
cursor = conn.cursor()
|
| 13 |
+
students = cursor.execute("SELECT name FROM students").fetchall()
|
| 14 |
+
conn.close()
|
| 15 |
+
return [{"name": student[0]} for student in students]
|
| 16 |
+
|
| 17 |
+
# Initialize the embedding model
|
| 18 |
+
model = SentenceTransformer('sentence-transformers/all-MiniLM-L6-v2')
|
| 19 |
+
|
| 20 |
+
# Create the query embedding
|
| 21 |
+
query_embedding = model.encode([query])
|
| 22 |
+
|
| 23 |
+
# Initialize ChromaDB client and get the collection
|
| 24 |
+
client = chromadb.PersistentClient(path="rag/vector_store")
|
| 25 |
+
collection = client.get_collection("university_data")
|
| 26 |
+
|
| 27 |
+
# Perform the search
|
| 28 |
+
results = collection.query(
|
| 29 |
+
query_embeddings=query_embedding,
|
| 30 |
+
n_results=top_k
|
| 31 |
+
)
|
| 32 |
+
|
| 33 |
+
return results['documents'][0]
|
| 34 |
+
|
| 35 |
+
if __name__ == '__main__':
|
| 36 |
+
# Example usage
|
| 37 |
+
test_query = "What courses are available?"
|
| 38 |
+
results = search(test_query)
|
| 39 |
+
print(f"Results for '{test_query}':")
|
| 40 |
+
for result in results:
|
| 41 |
+
print(result)
|
| 42 |
+
|
| 43 |
+
test_query_students = "Give me the names of the students"
|
| 44 |
+
results_students = search(test_query_students)
|
| 45 |
+
print(f"\nResults for '{test_query_students}':")
|
| 46 |
+
for student in results_students:
|
| 47 |
+
print(student['name'])
|