Samuel Oberhofer commited on
Commit
2322128
·
1 Parent(s): ff27905

feat: Implement retriever and guardrail modules

Browse files
Files changed (3) hide show
  1. guards/input.py +30 -0
  2. guards/output.py +27 -0
  3. 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'])