Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 78 additions & 0 deletions backend/src/modules/paper2figure/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,10 +49,31 @@ class MinerUConfig:
request_timeout: float = 120.0


@dataclass
class AI4ScholarConfig:
"""AI4Scholar paper-search service configuration."""

base_url: str = "https://ai4scholar.net/graph/v1"
api_key: str = ""


@dataclass
class RelatedPaperRetrievalConfig:
"""Runtime settings for dynamic related-paper retrieval."""

query_count: int = 3
topk: int = 10
search_results_per_query: int = 8
search_max_concurrency: int = 1
embedding_model: str = "text-embedding-3-small"
cache_db_path: str = "output/paper_cache.sqlite3"


NODE_CONFIGS: dict[str, NodeLLMConfig] = {
"text_classifier": NodeLLMConfig(model="gpt-4o-mini", temperature=0.0, max_tokens=1000),
"table_from_text": NodeLLMConfig(model="gpt-4o-mini", temperature=0.0, max_tokens=4000),
"paper_idea": NodeLLMConfig(model="gpt-4o-mini", temperature=0.1, max_tokens=4000),
"paper_query_planner": NodeLLMConfig(model="gpt-4o-mini", temperature=0.0, max_tokens=1200),
"output_decision": NodeLLMConfig(model="gpt-4o-mini", temperature=0.1, max_tokens=2000),
"chart_plan": NodeLLMConfig(model="gpt-4o-mini", temperature=0.1, max_tokens=3000),
"chart_code": NodeLLMConfig(model="gpt-4o", temperature=0.0, max_tokens=5000),
Expand Down Expand Up @@ -149,6 +170,55 @@ def get_mineru_config() -> MinerUConfig:
)


def get_ai4scholar_config() -> AI4ScholarConfig:
"""Return AI4Scholar configuration from YAML or environment."""

services = _load_module_config().get("services", {})
payload = services.get("ai4scholar", {}) if isinstance(services, dict) else {}
if not isinstance(payload, dict):
payload = {}

return AI4ScholarConfig(
base_url=str(
payload.get("base_url")
or payload.get("url")
or os.getenv("AI4SCHOLAR_BASE_URL")
or "https://ai4scholar.net/graph/v1"
).rstrip("/"),
api_key=str(
payload.get("api_key")
or os.getenv("AI4SCHOLAR_API_KEY")
or os.getenv("AI4SCHOLAR_KEY")
or ""
),
)


def get_related_paper_retrieval_config() -> RelatedPaperRetrievalConfig:
"""Return dynamic related-paper retrieval settings."""

payload = _load_module_config().get("paper_retrieval", {})
if not isinstance(payload, dict):
payload = {}

return RelatedPaperRetrievalConfig(
query_count=_positive_int(payload.get("query_count"), 3),
topk=_positive_int(payload.get("topk"), 10),
search_results_per_query=_positive_int(payload.get("search_results_per_query"), 8),
search_max_concurrency=_positive_int(payload.get("search_max_concurrency"), 1),
embedding_model=str(
payload.get("embedding_model")
or os.getenv("PAPER2FIGURE_EMBEDDING_MODEL")
or "text-embedding-3-small"
),
cache_db_path=str(
payload.get("cache_db_path")
or os.getenv("PAPER2FIGURE_PAPER_CACHE_DB_PATH")
or "output/paper_cache.sqlite3"
),
)


def create_node_model(node_name: str, model_override: str = "") -> ChatOpenAI:
"""Create LLM instance for a specific node."""

Expand All @@ -165,6 +235,14 @@ def create_node_model(node_name: str, model_override: str = "") -> ChatOpenAI:
return ChatOpenAI(**kwargs)


def _positive_int(value: Any, default: int) -> int:
try:
parsed = int(value)
except (TypeError, ValueError):
parsed = default
return max(1, parsed)


def _get_service_config(service_name: str) -> ServiceAPIConfig:
services = _load_module_config().get("services", {})
if not isinstance(services, dict):
Expand Down
21 changes: 21 additions & 0 deletions backend/src/modules/paper2figure/prompts.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,27 @@ def build_paper_idea_prompt(paper_content: str, language: str) -> tuple[str, str
return system, user


def build_paper_query_plan_prompt(target_text: str, query_count: int = 3) -> tuple[str, str]:
"""Prompt for planning concise academic paper-search queries."""

safe_count = max(1, int(query_count or 1))
system = f"""
You generate search queries for retrieving related academic papers.
Return valid JSON only with schema:
{{
"queries": ["query 1", "query 2"]
}}
Rules:
- Return exactly {safe_count} concise English search queries.
- Focus on the target paper problem setting, method family, task, and framework/pipeline concepts.
- Do not include generic words like survey, related work, or paper unless they are part of the topic.
- Prefer noun phrases that work well in academic search engines.
- Each query should be under 12 words.
""".strip()
user = f"Target paper text:\n{target_text[:12000]}"
return system, user


def build_output_decision_prompt(
text_input: str,
paper_idea_summary: str,
Expand Down
10 changes: 9 additions & 1 deletion backend/src/modules/paper2figure/utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,12 @@
parse_mineru_zip_url,
parse_mineru_zip_urls,
)
from .related_papers import get_topk_related_papers
from .reference_figures import collect_reference_figure_items
from .related_papers import (
collect_related_paper_debug,
get_topk_related_paper_records,
get_topk_related_papers,
)
from .table_utils import (
has_chartable_table,
render_table_preview,
Expand Down Expand Up @@ -56,6 +61,9 @@
"parse_mineru_zip_urls",
"extract_framework_image_items",
"collect_framework_image_items",
"collect_reference_figure_items",
"collect_related_paper_debug",
"get_topk_related_paper_records",
"get_topk_related_papers",
"has_chartable_table",
"render_table_preview",
Expand Down
185 changes: 185 additions & 0 deletions backend/src/modules/paper2figure/utils/paper_cache.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,185 @@
"""Simple SQLite cache for retrieved papers and embeddings."""

from __future__ import annotations

import json
import sqlite3
import time
from pathlib import Path
from typing import Any


def normalize_title_key(title: str) -> str:
return " ".join(str(title or "").strip().lower().split())


class PaperCache:
"""Persist paper metadata and embeddings by normalized title."""

def __init__(self, db_path: str):
self.db_path = Path(db_path).expanduser().resolve()
self.db_path.parent.mkdir(parents=True, exist_ok=True)
self._init_db()

def _connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(str(self.db_path))
connection.row_factory = sqlite3.Row
return connection

def _init_db(self) -> None:
with self._connect() as connection:
connection.execute(
"""
CREATE TABLE IF NOT EXISTS papers (
title_key TEXT PRIMARY KEY,
title TEXT NOT NULL,
payload_json TEXT NOT NULL,
embedding_model TEXT NOT NULL DEFAULT '',
embedding_text TEXT NOT NULL DEFAULT '',
embedding_json TEXT NOT NULL DEFAULT '',
updated_at REAL NOT NULL
)
"""
)
connection.commit()

def get_paper(self, title: str) -> dict[str, Any] | None:
title_key = normalize_title_key(title)
if not title_key:
return None
with self._connect() as connection:
row = connection.execute(
"SELECT payload_json FROM papers WHERE title_key = ?",
(title_key,),
).fetchone()
if row is None:
return None
try:
payload = json.loads(str(row["payload_json"] or "{}"))
except json.JSONDecodeError:
return None
return payload if isinstance(payload, dict) else None

def upsert_paper(self, paper: dict[str, Any]) -> None:
title = str(paper.get("title", "") or "").strip()
title_key = normalize_title_key(title)
if not title_key:
return
payload_json = json.dumps(paper, ensure_ascii=False)
with self._connect() as connection:
connection.execute(
"""
INSERT INTO papers (
title_key,
title,
payload_json,
embedding_model,
embedding_text,
embedding_json,
updated_at
)
VALUES (
?,
?,
?,
COALESCE((SELECT embedding_model FROM papers WHERE title_key = ?), ''),
COALESCE((SELECT embedding_text FROM papers WHERE title_key = ?), ''),
COALESCE((SELECT embedding_json FROM papers WHERE title_key = ?), ''),
?
)
ON CONFLICT(title_key) DO UPDATE SET
title = excluded.title,
payload_json = excluded.payload_json,
updated_at = excluded.updated_at
""",
(
title_key,
title,
payload_json,
title_key,
title_key,
title_key,
time.time(),
),
)
connection.commit()

def get_embedding(
self,
*,
title: str,
embedding_model: str,
embedding_text: str,
) -> list[float] | None:
title_key = normalize_title_key(title)
if not title_key:
return None
with self._connect() as connection:
row = connection.execute(
"""
SELECT embedding_json
FROM papers
WHERE title_key = ? AND embedding_model = ? AND embedding_text = ?
""",
(title_key, embedding_model, embedding_text),
).fetchone()
if row is None or not row["embedding_json"]:
return None
try:
payload = json.loads(str(row["embedding_json"]))
except json.JSONDecodeError:
return None
if not isinstance(payload, list):
return None
try:
return [float(value) for value in payload]
except (TypeError, ValueError):
return None

def upsert_embedding(
self,
*,
title: str,
paper: dict[str, Any],
embedding_model: str,
embedding_text: str,
embedding: list[float],
) -> None:
title_value = str(title or paper.get("title", "") or "").strip()
title_key = normalize_title_key(title_value)
if not title_key:
return
payload_json = json.dumps(paper, ensure_ascii=False)
embedding_json = json.dumps(embedding)
with self._connect() as connection:
connection.execute(
"""
INSERT INTO papers (
title_key,
title,
payload_json,
embedding_model,
embedding_text,
embedding_json,
updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(title_key) DO UPDATE SET
title = excluded.title,
payload_json = excluded.payload_json,
embedding_model = excluded.embedding_model,
embedding_text = excluded.embedding_text,
embedding_json = excluded.embedding_json,
updated_at = excluded.updated_at
""",
(
title_key,
title_value,
payload_json,
embedding_model,
embedding_text,
embedding_json,
time.time(),
),
)
connection.commit()
Loading