File size: 2,325 Bytes
abe7564
 
db774b8
 
 
 
abe7564
db774b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12a886e
db774b8
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
from typing import Optional
from neo4j import GraphDatabase, Driver
from src.models import Entity, Relationship
from src.config import NEO4J_URI, NEO4J_USERNAME, NEO4J_PASSWORD, NEO4J_DATABASE


_driver: Optional[Driver] = None


def _get_driver():
    global _driver
    if _driver is None:
        _driver = GraphDatabase.driver(NEO4J_URI, auth=(NEO4J_USERNAME, NEO4J_PASSWORD))
    return _driver


def create_entity(entity: Entity):
    with _get_driver().session(database=NEO4J_DATABASE) as session:
        session.run(
            "MERGE (e:Entity {name: $name}) "
            "SET e.type = $type, e.chunk_id = $chunk_id",
            name=entity.name, type=entity.type, chunk_id=entity.chunk_id,
        )


def create_relationship(rel: Relationship):
    with _get_driver().session(database=NEO4J_DATABASE) as session:
        session.run(
            "MATCH (a:Entity {name: $source}) "
            "MATCH (b:Entity {name: $target}) "
            "MERGE (a)-[r:RELATES_TO {type: $relation}]->(b) "
            "SET r.chunk_id = $chunk_id",
            source=rel.source_entity,
            target=rel.target_entity,
            relation=rel.relation_type,
            chunk_id=rel.chunk_id,
        )


def store_graph(entities: list[Entity], relationships: list[Relationship]):
    seen_names = set()
    for e in entities:
        if e.name not in seen_names:
            create_entity(e)
            seen_names.add(e.name)
    for r in relationships:
        create_relationship(r)


def get_connected_entities(entity_name: str, depth: int = 2) -> list[dict]:
    with _get_driver().session(database=NEO4J_DATABASE) as session:
        result = session.run(
            f"MATCH (e:Entity {{name: $name}})-[r*1..{depth}]-(connected) "
            "RETURN connected.name AS name, connected.type AS type, "
            "reduce(s = '', rel IN r | s + type(rel) + ' -> ') AS path",
            name=entity_name,
        )
        return [dict(r) for r in result]


def extract_entity_names() -> list[str]:
    with _get_driver().session(database=NEO4J_DATABASE) as session:
        result = session.run(
            "MATCH (e:Entity) RETURN e.name AS name LIMIT 200"
        )
        return [r["name"] for r in result]


def close():
    global _driver
    if _driver:
        _driver.close()
        _driver = None