jaothan/DockerGenAI_Streamlit
0
1import os
2import requests
3from dotenv import load_dotenv
4from langchain_community.graphs import Neo4jGraph
5import streamlit as st
6from streamlit.logger import get_logger
7from chains import load_embedding_model
8from utils import create_constraints, create_vector_index
9from PIL import Image
10
11load_dotenv(".env")
12
13url = os.getenv("NEO4J_URI")
14username = os.getenv("NEO4J_USERNAME")
15password = os.getenv("NEO4J_PASSWORD")
16ollama_base_url = os.getenv("OLLAMA_BASE_URL")
17embedding_model_name = os.getenv("EMBEDDING_MODEL")
18# Remapping for Langchain Neo4j integration
19os.environ["NEO4J_URL"] = url
20
21logger = get_logger(__name__)
22
23so_api_base_url = "https://api.stackexchange.com/2.3/search/advanced"
24
25embeddings, dimension = load_embedding_model(
26 embedding_model_name, config={"ollama_base_url": ollama_base_url}, logger=logger
27)
28
29# if Neo4j is local, you can go to http://localhost:7474/ to browse the database
30neo4j_graph = Neo4jGraph(
31 url=url, username=username, password=password, refresh_schema=False
32)
33
34create_constraints(neo4j_graph)
35create_vector_index(neo4j_graph)
36
37
38def load_so_data(tag: str = "neo4j", page: int = 1) -> None:
39 parameters = (
40 f"?pagesize=100&page={page}&order=desc&sort=creation&answers=1&tagged={tag}"
41 "&site=stackoverflow&filter=!*236eb_eL9rai)MOSNZ-6D3Q6ZKb0buI*IVotWaTb"
42 )
43 data = requests.get(so_api_base_url + parameters).json()
44 insert_so_data(data)
45
46
47def load_high_score_so_data() -> None:
48 parameters = (
49 f"?fromdate=1664150400&order=desc&sort=votes&site=stackoverflow&"
50 "filter=!.DK56VBPooplF.)bWW5iOX32Fh1lcCkw1b_Y6Zkb7YD8.ZMhrR5.FRRsR6Z1uK8*Z5wPaONvyII"
51 )
52 data = requests.get(so_api_base_url + parameters).json()
53 insert_so_data(data)
54
55
56def insert_so_data(data: dict) -> None:
57 # Calculate embedding values for questions and answers
58 for q in data["items"]:
59 question_text = q["title"] + "\n" + q["body_markdown"]
60 q["embedding"] = embeddings.embed_query(question_text)
61 for a in q["answers"]:
62 a["embedding"] = embeddings.embed_query(
63 question_text + "\n" + a["body_markdown"]
64 )
65
66 # Cypher, the query language of Neo4j, is used to import the data
67 # https://neo4j.com/docs/getting-started/cypher-intro/
68 # https://neo4j.com/docs/cypher-cheat-sheet/5/auradb-enterprise/
69 import_query = """
70 UNWIND $data AS q
71 MERGE (question:Question {id:q.question_id})
72 ON CREATE SET question.title = q.title, question.link = q.link, question.score = q.score,
73 question.favorite_count = q.favorite_count, question.creation_date = datetime({epochSeconds: q.creation_date}),
74 question.body = q.body_markdown, question.embedding = q.embedding
75 FOREACH (tagName IN q.tags |
76 MERGE (tag:Tag {name:tagName})
77 MERGE (question)-[:TAGGED]->(tag)
78 )
79 FOREACH (a IN q.answers |
80 MERGE (question)<-[:ANSWERS]-(answer:Answer {id:a.answer_id})
81 SET answer.is_accepted = a.is_accepted,
82 answer.score = a.score,
83 answer.creation_date = datetime({epochSeconds:a.creation_date}),
84 answer.body = a.body_markdown,
85 answer.embedding = a.embedding
86 MERGE (answerer:User {id:coalesce(a.owner.user_id, "deleted")})
87 ON CREATE SET answerer.display_name = a.owner.display_name,
88 answerer.reputation= a.owner.reputation
89 MERGE (answer)<-[:PROVIDED]-(answerer)
90 )
91 WITH * WHERE NOT q.owner.user_id IS NULL
92 MERGE (owner:User {id:q.owner.user_id})
93 ON CREATE SET owner.display_name = q.owner.display_name,
94 owner.reputation = q.owner.reputation
95 MERGE (owner)-[:ASKED]->(question)
96 """
97 neo4j_graph.query(import_query, {"data": data["items"]})
98
99
100# Streamlit
101def get_tag() -> str:
102 input_text = st.text_input(
103 "Which tag questions do you want to import?", value="neo4j"
104 )
105 return input_text
106
107
108def get_pages():
109 col1, col2 = st.columns(2)
110 with col1:
111 num_pages = st.number_input(
112 "Number of pages (100 questions per page)", step=1, min_value=1
113 )
114 with col2:
115 start_page = st.number_input("Start page", step=1, min_value=1)
116 st.caption("Only questions with answers will be imported.")
117 return (int(num_pages), int(start_page))
118
119
120def render_page():
121 datamodel_image = Image.open("./images/datamodel.png")
122 st.header("StackOverflow Loader")
123 st.subheader("Choose StackOverflow tags to load into Neo4j")
124 st.caption("Go to http://localhost:7474/ to explore the graph.")
125
126 user_input = get_tag()
127 num_pages, start_page = get_pages()
128
129 if st.button("Import", type="primary"):
130 with st.spinner("Loading... This might take a minute or two."):
131 try:
132 for page in range(1, num_pages + 1):
133 load_so_data(user_input, start_page + (page - 1))
134 st.success("Import successful", icon="โ
")
135 st.caption("Data model")
136 st.image(datamodel_image)
137 st.caption("Go to http://localhost:7474/ to interact with the database")
138 except Exception as e:
139 st.error(f"Error: {e}", icon="๐จ")
140 with st.expander("Highly ranked questions rather than tags?"):
141 if st.button("Import highly ranked questions"):
142 with st.spinner("Loading... This might take a minute or two."):
143 try:
144 load_high_score_so_data()
145 st.success("Import successful", icon="โ
")
146 except Exception as e:
147 st.error(f"Error: {e}", icon="๐จ")
148
149
150render_page()
151 