1010
1111from codegraphcontext .utils .debug_log import debug_log , info_logger , error_logger , warning_logger
1212
13+ class Neo4jDriverWrapper :
14+ """
15+ A simple wrapper around the Neo4j Driver to inject a database name into session() calls.
16+ """
17+ def __init__ (self , driver : Driver , database : str = None ):
18+ self ._driver = driver
19+ self ._database = database
20+
21+ def session (self , ** kwargs ):
22+ """Proxy method to get a session from the underlying driver."""
23+ if self ._database and 'database' not in kwargs :
24+ kwargs ["database" ] = self ._database
25+ return self ._driver .session (** kwargs )
26+
27+ def close (self ):
28+ """Proxy method to close the underlying driver."""
29+ self ._driver .close ()
30+
1331class DatabaseManager :
1432 """
1533 Manages the Neo4j database driver as a singleton to ensure only one
@@ -42,6 +60,7 @@ def __init__(self):
4260 self .neo4j_uri = os .getenv ('NEO4J_URI' )
4361 self .neo4j_username = os .getenv ('NEO4J_USERNAME' , 'neo4j' )
4462 self .neo4j_password = os .getenv ('NEO4J_PASSWORD' )
63+ self .neo4j_database = os .getenv ('NEO4J_DATABASE' ) # Optional, if not set, will use default database configured in Neo4j
4564 self ._initialized = True
4665
4766 def get_driver (self ) -> Driver :
@@ -53,7 +72,7 @@ def get_driver(self) -> Driver:
5372 ValueError: If Neo4j credentials are not set in environment variables.
5473
5574 Returns:
56- The active Neo4j Driver instance.
75+ The a wrapper for Neo4j Driver instance.
5776 """
5877 if self ._driver is None :
5978 with self ._lock :
@@ -100,7 +119,7 @@ def get_driver(self) -> Driver:
100119 self ._driver .close ()
101120 self ._driver = None
102121 raise
103- return self ._driver
122+ return Neo4jDriverWrapper ( self ._driver , database = self . neo4j_database )
104123
105124 def close_driver (self ):
106125 """Closes the Neo4j driver connection if it exists."""
@@ -116,7 +135,10 @@ def is_connected(self) -> bool:
116135 if self ._driver is None :
117136 return False
118137 try :
119- with self ._driver .session () as session :
138+ session_kwargs = {}
139+ if self .neo4j_database :
140+ session_kwargs ['database' ] = self .neo4j_database
141+ with self ._driver .session (** session_kwargs ) as session :
120142 session .run ("RETURN 1" ).consume ()
121143 return True
122144 except Exception :
@@ -163,7 +185,7 @@ def validate_config(uri: str, username: str, password: str) -> Tuple[bool, Optio
163185 return True , None
164186
165187 @staticmethod
166- def test_connection (uri : str , username : str , password : str ) -> Tuple [bool , Optional [str ]]:
188+ def test_connection (uri : str , username : str , password : str , database : str = None ) -> Tuple [bool , Optional [str ]]:
167189 """
168190 Tests the Neo4j database connection.
169191
@@ -206,7 +228,10 @@ def test_connection(uri: str, username: str, password: str) -> Tuple[bool, Optio
206228 # Now test Neo4j authentication
207229 driver = GraphDatabase .driver (uri , auth = (username , password ))
208230
209- with driver .session () as session :
231+ session_kwargs = {}
232+ if database :
233+ session_kwargs ['database' ] = database # Pass database to session if provided
234+ with driver .session (** session_kwargs ) as session :
210235 result = session .run ("RETURN 'Connection successful' as status" )
211236 result .single ()
212237
0 commit comments