genmod / application_neo4j /app_algorithms.py
pbordescnil's picture
Resolve search source using requested node label
39be855
Raw
History Blame
15.7 kB
# Importation des bibliothèques nécessaires
from graphdatascience import GraphDataScience
from typing import Dict, List, Any
import pandas as pd
from translations import t
from neo4j import Query
def run_gds_bfs(
gds: GraphDataScience,
natural_graph_name: str,
reverse_graph_name: str,
source_name: str,
max_depth: int = None,
expert=False,
source_labels=None,
job_id: str = None,
progress_callback=None,
neo4j_driver=None,
cancel_check=None,
) -> Dict[str, Any]:
"""
Parcourt les descendants et ascendants avec le fork OpenGDS.
Le plugin personnalisé ajoute au résultat standard ``nodeIds`` une liste
``depths`` alignée. La profondeur minimale de chaque nœud est ainsi calculée
pendant le BFS, sans second parcours Cypher.
Returns:
L'identifiant et le label de la source, ainsi que les résultats GDS
des parcours descendant et ascendant.
"""
requested_source_labels = [
label
for label in (source_labels or [])
if label in ("Model", "Dataset", "Author")
]
try:
source_id_result = gds.run_cypher(
"""
MATCH (n {name: $source_name})
WHERE size($source_labels) = 0
OR any(label IN labels(n) WHERE label IN $source_labels)
RETURN id(n) AS id, labels(n) AS label
""",
{
"source_name": source_name,
"source_labels": requested_source_labels,
},
)
if source_id_result.empty:
print(f"Le modèle ou dataset avec le nom '{source_name}' n'a pas été trouvé.")
return None
label_preferences = requested_source_labels or [
"Model",
"Dataset",
"Author",
]
selected_source = None
source_label = None
for preferred_label in label_preferences:
matching_sources = source_id_result[
source_id_result["label"].apply(
lambda node_labels: preferred_label in node_labels
)
]
if not matching_sources.empty:
selected_source = matching_sources.iloc[0]
source_label = preferred_label
break
if selected_source is None:
selected_source = source_id_result.iloc[0]
source_label = selected_source["label"][0]
if source_label == "Author" and not expert:
print(f"Le modèle ou dataset avec le nom '{source_name}' n'a pas été trouvé.")
return None
source_node_id = int(selected_source["id"])
except Exception as e:
print(f"Erreur lors de la recherche de l'ID du nœud source pour '{source_name}': {e}")
return None
bfs_params = {"sourceNode": source_node_id}
if max_depth is not None:
bfs_params["maxDepth"] = max_depth
def run_bfs(graph_name, params):
if cancel_check:
cancel_check()
if neo4j_driver is None:
graph = gds.graph.get(graph_name)
result = gds.bfs.stream(graph, **params)
else:
query = Query(
"""
CALL gds.bfs.stream($graph_name, $configuration)
YIELD nodeIds, depths
RETURN nodeIds, depths
""",
metadata={"search_job_id": job_id},
)
with neo4j_driver.session() as neo4j_session:
records = neo4j_session.run(
query,
{
"graph_name": graph_name,
"configuration": params,
},
)
result = pd.DataFrame(
record.data() for record in records
)
if cancel_check:
cancel_check()
return result
desc_job_id = f"{job_id}-descendants" if job_id else None
if desc_job_id:
bfs_params["jobId"] = desc_job_id
if progress_callback:
progress_callback("searching_descendants", desc_job_id)
desc_df = run_bfs(natural_graph_name, bfs_params)
_validate_depths_result(desc_df)
print("BFS descendants terminé avec les profondeurs.")
asc_job_id = f"{job_id}-ancestors" if job_id else None
if asc_job_id:
bfs_params["jobId"] = asc_job_id
elif "jobId" in bfs_params:
del bfs_params["jobId"]
if progress_callback:
progress_callback("searching_ancestors", asc_job_id)
asc_df = run_bfs(reverse_graph_name, bfs_params)
_validate_depths_result(asc_df)
print("BFS ascendants terminé avec les profondeurs.")
return {
"source_node": source_node_id,
"source_label": source_label,
"descendant": desc_df,
"ascendant": asc_df,
}
def _validate_depths_result(result: pd.DataFrame) -> None:
"""Fail explicitly if Neo4j did not load the custom OpenGDS plugin."""
if result.empty:
return
if "depths" not in result.columns:
raise RuntimeError(
"Le plugin OpenGDS personnalisé n'est pas chargé : "
"gds.bfs.stream ne renvoie pas la colonne depths."
)
node_ids = result["nodeIds"].iloc[0]
depths = result["depths"].iloc[0]
if len(node_ids) != len(depths):
raise RuntimeError(
"Résultat BFS invalide : nodeIds et depths n'ont pas la même taille."
)
def get_genealogy_highlights(gds: "GraphDataScience", model_name: str, num_highlights: int = 2, lang: str = "fr") -> Dict:
"""
Trouve les modèles clés dans l'ascendance et la descendance (1er/2e plus cités/téléchargés).
Args:
gds: L'instance de GraphDataScience.
model_name: Le nom du modèle de départ.
num_highlights: Le nombre de modèles à récupérer pour chaque catégorie (par défaut 2).
Returns:
Un dictionnaire contenant les listes de modèles unifiés pour l'affichage.
"""
highlights = {
"desc_unique_models": [],
"asc_unique_models": []
}
# --- DÉFINITION CENTRALE DES BADGES ---
# Centraliser les badges ici rend le code beaucoup plus facile à modifier.
# Les classes CSS sont directement des classes Bootstrap 5 pour simplifier le rendu dans le template HTML.
badges_info = {
'desc_cited_1': {
'text': t('badge.desc_cited_1.text', lang),
'class': 'bg-success',
'title': t('badge.desc_cited_1.title', lang)
},
'desc_cited_2': {
'text': t('badge.desc_cited_2.text', lang),
'class': 'bg-success bg-opacity-75',
'title': t('badge.desc_cited_2.title', lang)
},
'desc_downloaded_1': {
'text': t('badge.desc_downloaded_1.text', lang),
'class': 'beta',
'title': t('badge.desc_downloaded_1.title', lang)
},
'desc_downloaded_2': {
'text': t('badge.desc_downloaded_2.text', lang),
'class': 'alpha',
'title': t('badge.desc_downloaded_2.title', lang)
},
'asc_foundation': {
'text': t('badge.asc_foundation.text', lang),
'class': 'bg-warning text-dark',
'title': t('badge.asc_foundation.title', lang)
},
'asc_cited_1': {
'text': t('badge.asc_cited_1.text', lang),
'class': 'bg-success',
'title': t('badge.asc_cited_1.title', lang)
},
'asc_cited_2': {
'text': t('badge.asc_cited_2.text', lang),
'class': 'bg-success bg-opacity-75',
'title': t('badge.asc_cited_2.title', lang)
},
'asc_downloaded_1': {
'text': t('badge.asc_downloaded_1.text', lang),
'class': 'beta',
'title': t('badge.asc_downloaded_1.title', lang)
},
'asc_downloaded_2': {
'text': t('badge.asc_downloaded_2.text', lang),
'class': 'alpha',
'title': t('badge.asc_downloaded_2.title', lang)
},
}
def process_and_assign_badges(
unified_dict: Dict,
model_list: List[Dict],
badge_keys: List[str]
):
"""
Fonction utilitaire pour ajouter des modèles et leurs badges à un dictionnaire unifié.
Cela évite la duplication de code pour chaque catégorie (cité, téléchargé, etc.).
"""
for i, model in enumerate(model_list):
if i < len(badge_keys): # S'assurer qu'on a un badge défini pour ce rang
model_name_key = model['name']
badge_key = badge_keys[i]
# Ajouter le modèle au dictionnaire s'il n'y est pas déjà
if model_name_key not in unified_dict:
unified_dict[model_name_key] = model.copy()
unified_dict[model_name_key]['badges'] = []
# Ajouter le badge correspondant
badge_to_add = badges_info[badge_key]
if badge_to_add not in unified_dict[model_name_key]['badges']:
unified_dict[model_name_key]['badges'].append(badge_to_add)
# ==========================================================================
# 1. GESTION DE LA DESCENDANCE
# ==========================================================================
desc_downloads_query = """
MATCH (start:Model {name: $model_name})-[:USED_IN*1..]->(descendant:Model)
WHERE start <> descendant
WITH descendant, size([(m:Model)<-[:USED_IN]-(descendant) | m]) AS citation_count
RETURN descendant.name AS name, citation_count, descendant.downloads AS downloads, descendant.task AS task, descendant.license AS license, descendant.likes AS likes, descendant.createdAt AS createdAt
ORDER BY descendant.downloads DESC, descendant.name ASC
LIMIT $limit
"""
desc_cited_query = """
MATCH (start:Model {name: $model_name})-[:USED_IN*1..]->(descendant:Model)
WHERE start <> descendant
WITH descendant, size([(m:Model)<-[:USED_IN]-(descendant) | m]) AS citation_count
RETURN descendant.name AS name, citation_count, descendant.task AS task, descendant.downloads AS downloads, descendant.license AS license, descendant.likes AS likes, descendant.createdAt AS createdAt
ORDER BY citation_count DESC, descendant.name ASC
LIMIT $limit
"""
try:
params = {"model_name": model_name, "limit": num_highlights}
desc_downloaded_list = gds.run_cypher(desc_downloads_query, params).to_dict('records')
desc_cited_list = gds.run_cypher(desc_cited_query, params).to_dict('records')
desc_unified_models = {}
process_and_assign_badges(desc_unified_models, desc_cited_list, ['desc_cited_1', 'desc_cited_2'])
process_and_assign_badges(desc_unified_models, desc_downloaded_list, ['desc_downloaded_1', 'desc_downloaded_2'])
highlights["desc_unique_models"] = list(desc_unified_models.values())
except Exception as e:
print(f"Erreur lors de la recherche des descendants: {e}")
# ==========================================================================
# 2. GESTION DE L'ASCENDANCE
# ==========================================================================
asc_downloads_query = """
MATCH (ascendant:Model)-[:USED_IN*1..]->(start:Model {name: $model_name})
WHERE start <> ascendant
WITH ascendant, size([(m:Model)<-[:USED_IN]-(ascendant) | m]) AS citation_count
RETURN ascendant.name AS name, citation_count, ascendant.downloads AS downloads, ascendant.task AS task,
ascendant.license AS license, ascendant.likes AS likes, ascendant.createdAt AS createdAt
ORDER BY ascendant.downloads DESC
LIMIT 1 // On ne veut que LE plus téléchargé
"""
asc_cited_query = """
MATCH (ascendant:Model)-[:USED_IN*1..]->(start:Model {name: $model_name})
WHERE start <> ascendant
WITH ascendant, size([(m:Model)<-[:USED_IN]-(ascendant) | m]) AS citation_count
RETURN ascendant.name AS name, citation_count, ascendant.downloads AS downloads, ascendant.task AS task,
ascendant.license AS license, ascendant.likes AS likes, ascendant.createdAt AS createdAt
ORDER BY citation_count DESC
LIMIT 1 // On ne veut que LE plus cité
"""
foundation_query = """
MATCH (foundation:Model)-[:USED_IN*1..]->(start:Model {name: $model_name})
WHERE NOT EXISTS( (:Model)-[:USED_IN]->(foundation) )
WITH foundation, size([(m:Model)<-[:USED_IN]-(foundation) | m]) AS citation_count
RETURN DISTINCT foundation.name AS name, citation_count, foundation.downloads AS downloads, foundation.task AS task,
foundation.license AS license, foundation.likes AS likes, foundation.createdAt AS createdAt
LIMIT $limit
"""
try:
params = {"model_name": model_name, "limit": num_highlights}
asc_foundation_list = gds.run_cypher(foundation_query, params).to_dict('records')
asc_downloaded_list = gds.run_cypher(asc_downloads_query, params).to_dict('records')
asc_cited_list = gds.run_cypher(asc_cited_query, params).to_dict('records')
asc_unified_models = {}
# Ordre de priorité : Racine > Cité > Téléchargé
process_and_assign_badges(asc_unified_models, asc_foundation_list, ['asc_foundation'] * num_highlights) # Le badge racine s'applique à tous
process_and_assign_badges(asc_unified_models, asc_cited_list, ['asc_cited_1', 'asc_cited_2'])
process_and_assign_badges(asc_unified_models, asc_downloaded_list, ['asc_downloaded_1', 'asc_downloaded_2'])
highlights["asc_unique_models"] = list(asc_unified_models.values())
except Exception as e:
print(f"Erreur lors de la recherche des ascendants: {e}")
return highlights
def create_node_data(node_props, label):
"""
Construit un dictionnaire de données pour chaque noeud
à afficher dans le graphe front-end.
"""
base_data = {
"id": node_props.get("name", "")
}
if label == "Author":
return {
**base_data,
"label": node_props.get("type", "Unknown"),
"followers": node_props.get("followers", 1)
}
elif label == "Model":
licens_ =str(node_props.get("license", t("node.unknown", "fr"))).strip("[]")
if licens_ =="\'other\'" or pd.isna(licens_) or licens_ =="nan":
licens_ = t("node.other", "fr")
tache = node_props.get("task", "")
if tache =="unknown":
tache = t("node.unknown", "fr")
return {
**base_data,
"label": "Modèle",
"downloads": node_props.get("downloads", 1),
"likes": node_props.get("likes", 0),
"license": licens_,
"createdAt": node_props.get("createdAt", "inconnue"),
"createdAt_dataset": node_props.get("createdAt_dataset", "inconnue"),
"task": tache,
"author": node_props.get("author", ""),"dataset": node_props.get("dataset", ""),
"ascendantsCount": node_props.get("ascendantsCount", 0),"descendantsCount": node_props.get("descendantsCount", 0),
"citationCount": node_props.get("citationCount", 0), "distance":node_props.get("distance", 0)
}
else: # Dataset or other
return {
**base_data,
"label": "Dataset",
"downloads": node_props.get("downloads", 1),
"createdAt_dataset": node_props.get("createdAt_dataset", "inconnue")
}
return { "id": node_props['name'], "label": label, **node_props }