langchain-diffbot 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,297 @@
1
+ """Diffbot retrievers — Knowledge Graph and Web Search."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Callable, Sequence
6
+ from typing import Any
7
+
8
+ from langchain_core.callbacks import (
9
+ AsyncCallbackManagerForRetrieverRun,
10
+ CallbackManagerForRetrieverRun,
11
+ )
12
+ from langchain_core.documents import Document
13
+ from langchain_core.retrievers import BaseRetriever
14
+ from pydantic import Field
15
+
16
+ from langchain_diffbot._base import _BaseDiffbotComponent
17
+
18
+ DEFAULT_KG_CONTENT_FIELDS: tuple[str, ...] = ("description", "summary", "name")
19
+ """Default ordered priority for selecting `page_content` from a KG entity."""
20
+
21
+ DEFAULT_WEB_CONTENT_FIELDS: tuple[str, ...] = ("content", "snippet")
22
+ """Default ordered priority for selecting `page_content` from a web search result."""
23
+
24
+ DocumentMapper = Callable[[dict[str, Any]], Document]
25
+
26
+
27
+ def _dict_to_document(
28
+ source: dict[str, Any],
29
+ *,
30
+ content_fields: Sequence[str],
31
+ fields: Sequence[str] | None,
32
+ ) -> Document:
33
+ """Map a flat dict to a LangChain `Document` using a content-field priority list.
34
+
35
+ `page_content` is the first non-empty value among `content_fields`.
36
+ `metadata` is the remaining top-level keys, optionally narrowed to
37
+ `fields` (the projection allowlist).
38
+ """
39
+ page_content = ""
40
+ content_field_used: str | None = None
41
+ for f in content_fields:
42
+ value = source.get(f)
43
+ if value:
44
+ page_content = value if isinstance(value, str) else str(value)
45
+ content_field_used = f
46
+ break
47
+
48
+ if fields is None:
49
+ metadata = {k: v for k, v in source.items() if k != content_field_used}
50
+ else:
51
+ metadata = {
52
+ k: source[k] for k in fields if k in source and k != content_field_used
53
+ }
54
+ return Document(page_content=page_content, metadata=metadata)
55
+
56
+
57
+ def _resolve_k(default_k: int, kwargs: dict[str, Any]) -> int:
58
+ k = kwargs.get("k", default_k)
59
+ if not isinstance(k, int) or k <= 0:
60
+ msg = f"`k` must be a positive integer, got {k!r}."
61
+ raise ValueError(msg)
62
+ return k
63
+
64
+
65
+ class DiffbotKnowledgeGraphRetriever(_BaseDiffbotComponent, BaseRetriever):
66
+ """Retriever backed by the Diffbot Knowledge Graph DQL endpoint.
67
+
68
+ The `query` passed to `invoke` is a
69
+ [DQL](https://docs.diffbot.com/reference/dql-quickstart) expression
70
+ (e.g. `type:Organization industries:"Artificial Intelligence"`).
71
+
72
+ Example:
73
+ ```python
74
+ from langchain_diffbot import DiffbotKnowledgeGraphRetriever
75
+
76
+ retriever = DiffbotKnowledgeGraphRetriever(k=5)
77
+ retriever.invoke('type:Organization location.city.name:"Boston"')
78
+ ```
79
+
80
+ Shaping the output (recommended for agent/tool use, where large entity
81
+ payloads can blow past LLM input-token limits):
82
+
83
+ ```python
84
+ retriever = DiffbotKnowledgeGraphRetriever(
85
+ k=5,
86
+ fields=["id", "type", "name", "homepageUri", "nbEmployees"],
87
+ )
88
+ ```
89
+
90
+ For full control, pass `document_mapper`:
91
+
92
+ ```python
93
+ def mapper(entity):
94
+ return Document(
95
+ page_content=entity.get("summary", ""),
96
+ metadata={"id": entity["id"], "name": entity["name"]},
97
+ )
98
+
99
+ retriever = DiffbotKnowledgeGraphRetriever(document_mapper=mapper)
100
+ ```
101
+
102
+ For full SDK control, supply a pre-built client:
103
+
104
+ ```python
105
+ from diffbot import Diffbot
106
+ retriever = DiffbotKnowledgeGraphRetriever(
107
+ client=Diffbot(token=..., timeout=60.0),
108
+ )
109
+ ```
110
+ """
111
+
112
+ k: int = 10
113
+ """Default number of results. Can be overridden per call via `invoke(..., k=N)`."""
114
+
115
+ from_: int = 0
116
+ """Result offset (passes through to `dql(from_=...)`)."""
117
+
118
+ filter: str | None = None
119
+ """DQL filter expression (passes through to `dql(filter=...)`)."""
120
+
121
+ exportspec: str | None = None
122
+ """Export spec for non-JSON formats (passes through to `dql(exportspec=...)`)."""
123
+
124
+ format: str = "json"
125
+ """Response format. Defaults to `json`.
126
+
127
+ Other values return raw bytes (see `dql(raw=True)`).
128
+ """
129
+
130
+ extra: dict[str, str] | None = None
131
+ """Extra query params merged into the DQL request.
132
+
133
+ Passes through to `dql(extra=...)`.
134
+ """
135
+
136
+ fields: list[str] | None = None
137
+ """Allowlist of top-level entity keys to keep in `metadata`.
138
+
139
+ `None` (default) keeps every field. Set this to a small list like
140
+ `["id", "type", "name", "homepageUri"]` to drastically shrink Document
141
+ payloads — important when the retriever feeds an LLM tool call, since
142
+ full Diffbot KG entities can run thousands of tokens each.
143
+
144
+ Ignored when `document_mapper` is set.
145
+ """
146
+
147
+ content_fields: list[str] = Field(
148
+ default_factory=lambda: list(DEFAULT_KG_CONTENT_FIELDS)
149
+ )
150
+ """Ordered priority for selecting `page_content` from an entity.
151
+
152
+ The first key in this list with a non-empty value wins. The chosen key
153
+ is excluded from `metadata` to avoid duplicating data.
154
+
155
+ Ignored when `document_mapper` is set.
156
+ """
157
+
158
+ document_mapper: DocumentMapper | None = None
159
+ """Optional override mapping a raw entity dict to a `Document`."""
160
+
161
+ def _hit_to_document(self, hit: dict[str, Any]) -> Document:
162
+ # Diffbot returns each result as
163
+ # {"score": ..., "entity": {...}, "entity_ctx": ...}. Older shapes
164
+ # (and our tests) sometimes embed entity fields at the top level, so
165
+ # fall back to the hit itself when there's no nested entity.
166
+ entity = hit.get("entity", hit)
167
+ if self.document_mapper is not None:
168
+ return self.document_mapper(entity)
169
+ doc = _dict_to_document(
170
+ entity,
171
+ content_fields=self.content_fields,
172
+ fields=self.fields,
173
+ )
174
+ if "score" in hit:
175
+ doc.metadata["score"] = hit["score"]
176
+ return doc
177
+
178
+ def _dql_kwargs(self, size: int) -> dict[str, Any]:
179
+ return {
180
+ "size": size,
181
+ "from_": self.from_,
182
+ "format": self.format,
183
+ "filter": self.filter,
184
+ "exportspec": self.exportspec,
185
+ "extra": self.extra,
186
+ }
187
+
188
+ def _get_relevant_documents(
189
+ self,
190
+ query: str,
191
+ *,
192
+ run_manager: CallbackManagerForRetrieverRun,
193
+ **kwargs: Any,
194
+ ) -> list[Document]:
195
+ k = _resolve_k(self.k, kwargs)
196
+ with self._sync_db() as db:
197
+ body = db.dql(query, **self._dql_kwargs(k))
198
+ if not isinstance(body, dict):
199
+ msg = (
200
+ "DQL returned a non-JSON body; "
201
+ "set `format='json'` to use this retriever."
202
+ )
203
+ raise TypeError(msg)
204
+ return [self._hit_to_document(h) for h in body.get("data", [])[:k]]
205
+
206
+ async def _aget_relevant_documents(
207
+ self,
208
+ query: str,
209
+ *,
210
+ run_manager: AsyncCallbackManagerForRetrieverRun,
211
+ **kwargs: Any,
212
+ ) -> list[Document]:
213
+ k = _resolve_k(self.k, kwargs)
214
+ async with self._async_db() as db:
215
+ body = await db.dql(query, **self._dql_kwargs(k))
216
+ if not isinstance(body, dict):
217
+ msg = (
218
+ "DQL returned a non-JSON body; "
219
+ "set `format='json'` to use this retriever."
220
+ )
221
+ raise TypeError(msg)
222
+ return [self._hit_to_document(h) for h in body.get("data", [])[:k]]
223
+
224
+
225
+ class DiffbotWebSearchRetriever(_BaseDiffbotComponent, BaseRetriever):
226
+ """Retriever backed by Diffbot's web search API.
227
+
228
+ The `query` passed to `invoke` is a natural-language search string.
229
+ Results come back as `Document`s whose `page_content` is the page content
230
+ (or snippet) returned by Diffbot, with `title`, `pageUrl`, and `score` in
231
+ `metadata` by default.
232
+
233
+ Example:
234
+ ```python
235
+ from langchain_diffbot import DiffbotWebSearchRetriever
236
+
237
+ retriever = DiffbotWebSearchRetriever(k=5)
238
+ retriever.invoke("diffbot knowledge graph")
239
+ ```
240
+ """
241
+
242
+ k: int = 10
243
+ """Number of results to fetch. Maps to the SDK's `num_results`."""
244
+
245
+ max_tokens: int | None = None
246
+ """Optional cap on total content tokens.
247
+
248
+ Passes through to `web_search(max_tokens=...)`.
249
+ """
250
+
251
+ fields: list[str] | None = None
252
+ """Allowlist of result keys to keep in `metadata`.
253
+
254
+ `None` keeps every field. Ignored when `document_mapper` is set.
255
+ """
256
+
257
+ content_fields: list[str] = Field(
258
+ default_factory=lambda: list(DEFAULT_WEB_CONTENT_FIELDS)
259
+ )
260
+ """Ordered priority for selecting `page_content`. First non-empty wins.
261
+
262
+ Ignored when `document_mapper` is set.
263
+ """
264
+
265
+ document_mapper: DocumentMapper | None = None
266
+ """Optional override mapping a raw search result dict to a `Document`."""
267
+
268
+ def _hit_to_document(self, hit: dict[str, Any]) -> Document:
269
+ if self.document_mapper is not None:
270
+ return self.document_mapper(hit)
271
+ return _dict_to_document(
272
+ hit, content_fields=self.content_fields, fields=self.fields
273
+ )
274
+
275
+ def _get_relevant_documents(
276
+ self,
277
+ query: str,
278
+ *,
279
+ run_manager: CallbackManagerForRetrieverRun,
280
+ **kwargs: Any,
281
+ ) -> list[Document]:
282
+ k = _resolve_k(self.k, kwargs)
283
+ with self._sync_db() as db:
284
+ body = db.web_search(query, num_results=k, max_tokens=self.max_tokens)
285
+ return [self._hit_to_document(h) for h in body.get("search_results", [])[:k]]
286
+
287
+ async def _aget_relevant_documents(
288
+ self,
289
+ query: str,
290
+ *,
291
+ run_manager: AsyncCallbackManagerForRetrieverRun,
292
+ **kwargs: Any,
293
+ ) -> list[Document]:
294
+ k = _resolve_k(self.k, kwargs)
295
+ async with self._async_db() as db:
296
+ body = await db.web_search(query, num_results=k, max_tokens=self.max_tokens)
297
+ return [self._hit_to_document(h) for h in body.get("search_results", [])[:k]]