graph-knowledge-doc-parser 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.
- graph_knowledge_doc_parser-0.1.0.dist-info/METADATA +326 -0
- graph_knowledge_doc_parser-0.1.0.dist-info/RECORD +38 -0
- graph_knowledge_doc_parser-0.1.0.dist-info/WHEEL +4 -0
- graph_knowledge_doc_parser-0.1.0.dist-info/entry_points.txt +3 -0
- kg_doc_parser/__init__.py +9 -0
- kg_doc_parser/cast_hinting.py +19 -0
- kg_doc_parser/document_ingester_logger.py +766 -0
- kg_doc_parser/models.py +277 -0
- kg_doc_parser/ocr.py +752 -0
- kg_doc_parser/pdf2png.py +286 -0
- kg_doc_parser/semantic_document_splitting_layerwise_edits.py +3302 -0
- kg_doc_parser/text_processing_utils.py +30 -0
- kg_doc_parser/utils/__init__.py +0 -0
- kg_doc_parser/utils/bounded_threadpool_executor.py +37 -0
- kg_doc_parser/utils/file_loaders.py +405 -0
- kg_doc_parser/utils/langchain.py +220 -0
- kg_doc_parser/utils/log.py +135 -0
- kg_doc_parser/utils/version_chaining.py +1278 -0
- kg_doc_parser/workflow_ingest/__init__.py +187 -0
- kg_doc_parser/workflow_ingest/_kogwistar.py +13 -0
- kg_doc_parser/workflow_ingest/adapters.py +212 -0
- kg_doc_parser/workflow_ingest/cache.py +63 -0
- kg_doc_parser/workflow_ingest/cli.py +324 -0
- kg_doc_parser/workflow_ingest/clients.py +444 -0
- kg_doc_parser/workflow_ingest/demo_harness.py +427 -0
- kg_doc_parser/workflow_ingest/design.py +208 -0
- kg_doc_parser/workflow_ingest/handlers.py +617 -0
- kg_doc_parser/workflow_ingest/models.py +575 -0
- kg_doc_parser/workflow_ingest/ocr_pipeline.py +1581 -0
- kg_doc_parser/workflow_ingest/page_index.py +473 -0
- kg_doc_parser/workflow_ingest/parser_core.py +862 -0
- kg_doc_parser/workflow_ingest/parsing.py +249 -0
- kg_doc_parser/workflow_ingest/probe.py +164 -0
- kg_doc_parser/workflow_ingest/providers.py +412 -0
- kg_doc_parser/workflow_ingest/runners.py +546 -0
- kg_doc_parser/workflow_ingest/semantics.py +231 -0
- kg_doc_parser/workflow_ingest/service.py +112 -0
- kg_doc_parser/workflow_ingest/smoke_assets.py +62 -0
|
@@ -0,0 +1,412 @@
|
|
|
1
|
+
"""Provider-neutral OCR, parser, and embedding adapters.
|
|
2
|
+
|
|
3
|
+
This module keeps vendor-specific imports behind small factory functions so the
|
|
4
|
+
workflow code can stay neutral. The concrete vendor is selected by config, not
|
|
5
|
+
by the caller.
|
|
6
|
+
|
|
7
|
+
Quick examples
|
|
8
|
+
--------------
|
|
9
|
+
- OCR with Google GenAI:
|
|
10
|
+
- KG_DOC_OCR_PROVIDER=gemini
|
|
11
|
+
- KG_DOC_OCR_MODEL=gemini-2.5-flash
|
|
12
|
+
|
|
13
|
+
- OCR with a local Ollama vision model:
|
|
14
|
+
- KG_DOC_OCR_PROVIDER=ollama
|
|
15
|
+
- KG_DOC_OCR_MODEL=llava:latest
|
|
16
|
+
- KG_DOC_OCR_BASE_URL=http://127.0.0.1:11434
|
|
17
|
+
|
|
18
|
+
- Parser/LLM with OpenAI Chat Completions:
|
|
19
|
+
- KG_DOC_PARSER_PROVIDER=openai
|
|
20
|
+
- KG_DOC_PARSER_MODEL=gpt-4.1-mini
|
|
21
|
+
- KG_DOC_PARSER_API_KEY_ENV=OPENAI_API_KEY
|
|
22
|
+
|
|
23
|
+
- Parser/LLM with Google Vertex AI:
|
|
24
|
+
- KG_DOC_PARSER_PROVIDER=vertex
|
|
25
|
+
- KG_DOC_PARSER_MODEL=gemini-2.5-pro
|
|
26
|
+
- KG_DOC_PARSER_PROJECT=my-project
|
|
27
|
+
- KG_DOC_PARSER_LOCATION=us-central1
|
|
28
|
+
|
|
29
|
+
- Parser/LLM with Ollama:
|
|
30
|
+
- KG_DOC_PARSER_PROVIDER=ollama
|
|
31
|
+
- KG_DOC_PARSER_MODEL=llama3.1
|
|
32
|
+
- KG_DOC_PARSER_BASE_URL=http://127.0.0.1:11434
|
|
33
|
+
|
|
34
|
+
- Embeddings with a fake deterministic function for CI:
|
|
35
|
+
- KG_DOC_EMBED_PROVIDER=fake
|
|
36
|
+
- KG_DOC_EMBED_MODEL=kg-doc-parser-workflow-embedding-v1
|
|
37
|
+
|
|
38
|
+
Cookbook example
|
|
39
|
+
----------------
|
|
40
|
+
If you are parsing a cooking recipe, you can keep OCR on Gemini but route the
|
|
41
|
+
parser to OpenAI or Ollama:
|
|
42
|
+
|
|
43
|
+
settings = WorkflowProviderSettings(
|
|
44
|
+
ocr=ProviderEndpointConfig(provider="gemini", model="gemini-2.5-flash"),
|
|
45
|
+
parser=ProviderEndpointConfig(
|
|
46
|
+
provider="openai",
|
|
47
|
+
model="gpt-4.1-mini",
|
|
48
|
+
api_key_env="OPENAI_API_KEY",
|
|
49
|
+
),
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
That means the OCR step extracts the page text, and the parser step can then
|
|
53
|
+
turn the recipe into structured fields such as ingredients, tools, actions,
|
|
54
|
+
and inferred sections without changing workflow orchestration.
|
|
55
|
+
"""
|
|
56
|
+
|
|
57
|
+
from __future__ import annotations
|
|
58
|
+
|
|
59
|
+
import os
|
|
60
|
+
from dataclasses import dataclass
|
|
61
|
+
from typing import Annotated, Any, Callable, ClassVar, Literal, Optional, Protocol, Union, runtime_checkable, get_args, get_origin
|
|
62
|
+
|
|
63
|
+
from pydantic import BaseModel, Field
|
|
64
|
+
from pydantic_core import PydanticUndefined
|
|
65
|
+
from pydantic_extension.model_slicing import BackendField, FrontendField
|
|
66
|
+
from pydantic_extension.model_slicing.mixin import DtoField, ExcludeMode, LLMField, ModeSlicingMixin
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
class _FakeStructuredResponse:
|
|
70
|
+
def __init__(self, schema, payload: dict[str, Any]):
|
|
71
|
+
self.schema = schema
|
|
72
|
+
self.payload = payload
|
|
73
|
+
|
|
74
|
+
def invoke(self, messages, config=None):
|
|
75
|
+
parsed = self.schema.model_validate(self.payload)
|
|
76
|
+
return {"parsed": parsed, "raw": None, "parsing_error": None}
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class FakeChatModel:
|
|
80
|
+
"""Minimal structured-output compatible chat model for tests."""
|
|
81
|
+
|
|
82
|
+
def __init__(self, *, payload_factory: Callable[[Any], dict[str, Any]] | None = None) -> None:
|
|
83
|
+
self.payload_factory = payload_factory or _default_schema_payload
|
|
84
|
+
|
|
85
|
+
def with_structured_output(self, schema, include_raw: bool = True):
|
|
86
|
+
payload = self.payload_factory(schema)
|
|
87
|
+
return _FakeStructuredResponse(schema, payload)
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def _default_schema_payload(schema) -> dict[str, Any]:
|
|
91
|
+
def _value_for_field(field) -> Any:
|
|
92
|
+
annotation = getattr(field, "annotation", None)
|
|
93
|
+
origin = get_origin(annotation)
|
|
94
|
+
args = get_args(annotation)
|
|
95
|
+
if annotation is str:
|
|
96
|
+
return ""
|
|
97
|
+
if annotation is bool:
|
|
98
|
+
return False
|
|
99
|
+
if annotation is int:
|
|
100
|
+
return 0
|
|
101
|
+
if annotation is float:
|
|
102
|
+
return 0.0
|
|
103
|
+
if origin is list or annotation is list:
|
|
104
|
+
return []
|
|
105
|
+
if origin is dict or annotation is dict:
|
|
106
|
+
return {}
|
|
107
|
+
if origin is tuple:
|
|
108
|
+
return []
|
|
109
|
+
if origin is Literal and args:
|
|
110
|
+
return args[0]
|
|
111
|
+
if origin is Union and type(None) in args:
|
|
112
|
+
return None
|
|
113
|
+
if hasattr(annotation, "model_fields"):
|
|
114
|
+
return _default_schema_payload(annotation)
|
|
115
|
+
default = getattr(field, "default", PydanticUndefined)
|
|
116
|
+
if default is not PydanticUndefined and default is not None:
|
|
117
|
+
return default
|
|
118
|
+
return None
|
|
119
|
+
|
|
120
|
+
payload: dict[str, Any] = {}
|
|
121
|
+
for name, field in getattr(schema, "model_fields", {}).items():
|
|
122
|
+
value = _value_for_field(field)
|
|
123
|
+
if value is not None:
|
|
124
|
+
payload[name] = value
|
|
125
|
+
return payload
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
@runtime_checkable
|
|
129
|
+
class ChatModelProvider(Protocol):
|
|
130
|
+
def build(self, *, callbacks: list[Any] | None = None) -> Any: ...
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
@runtime_checkable
|
|
134
|
+
class EmbeddingFunctionProvider(Protocol):
|
|
135
|
+
def build(self) -> Callable[[list[str]], list[list[float]]]: ...
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
class ProviderEndpointConfig(ModeSlicingMixin, BaseModel):
|
|
139
|
+
default_include_modes: ClassVar[set[str]] = {"dto", "backend", "frontend", "llm"}
|
|
140
|
+
include_unmarked_for_modes: ClassVar[set[str]] = {"dto", "backend", "frontend", "llm"}
|
|
141
|
+
|
|
142
|
+
provider: Annotated[
|
|
143
|
+
Literal["gemini", "ollama", "openai", "vertex", "fake"],
|
|
144
|
+
DtoField(),
|
|
145
|
+
BackendField(),
|
|
146
|
+
FrontendField(),
|
|
147
|
+
LLMField(),
|
|
148
|
+
] = "gemini"
|
|
149
|
+
model: Annotated[str, DtoField(), BackendField(), FrontendField(), LLMField()] = "gemini-2.5-flash"
|
|
150
|
+
temperature: Annotated[float, DtoField(), BackendField(), FrontendField(), LLMField()] = 0.1
|
|
151
|
+
base_url: Annotated[
|
|
152
|
+
Optional[str],
|
|
153
|
+
DtoField(),
|
|
154
|
+
BackendField(),
|
|
155
|
+
FrontendField(),
|
|
156
|
+
ExcludeMode("llm"),
|
|
157
|
+
] = None
|
|
158
|
+
api_key_env: Annotated[
|
|
159
|
+
Optional[str],
|
|
160
|
+
DtoField(),
|
|
161
|
+
BackendField(),
|
|
162
|
+
FrontendField(),
|
|
163
|
+
ExcludeMode("llm"),
|
|
164
|
+
] = None
|
|
165
|
+
project: Annotated[
|
|
166
|
+
Optional[str],
|
|
167
|
+
DtoField(),
|
|
168
|
+
BackendField(),
|
|
169
|
+
FrontendField(),
|
|
170
|
+
ExcludeMode("llm"),
|
|
171
|
+
] = None
|
|
172
|
+
location: Annotated[
|
|
173
|
+
Optional[str],
|
|
174
|
+
DtoField(),
|
|
175
|
+
BackendField(),
|
|
176
|
+
FrontendField(),
|
|
177
|
+
ExcludeMode("llm"),
|
|
178
|
+
] = None
|
|
179
|
+
max_retries: Annotated[int, DtoField(), BackendField(), FrontendField(), LLMField()] = 2
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
class EmbeddingProviderConfig(ModeSlicingMixin, BaseModel):
|
|
183
|
+
default_include_modes: ClassVar[set[str]] = {"dto", "backend", "frontend", "llm"}
|
|
184
|
+
include_unmarked_for_modes: ClassVar[set[str]] = {"dto", "backend", "frontend", "llm"}
|
|
185
|
+
|
|
186
|
+
provider: Annotated[
|
|
187
|
+
Literal["fake", "openai", "vertex", "ollama"],
|
|
188
|
+
DtoField(),
|
|
189
|
+
BackendField(),
|
|
190
|
+
FrontendField(),
|
|
191
|
+
LLMField(),
|
|
192
|
+
] = "fake"
|
|
193
|
+
model: Annotated[str, DtoField(), BackendField(), FrontendField(), LLMField()] = "kg-doc-parser-workflow-embedding-v1"
|
|
194
|
+
dimension: Annotated[int, DtoField(), BackendField(), FrontendField(), LLMField()] = 2
|
|
195
|
+
base_url: Annotated[
|
|
196
|
+
Optional[str],
|
|
197
|
+
DtoField(),
|
|
198
|
+
BackendField(),
|
|
199
|
+
FrontendField(),
|
|
200
|
+
ExcludeMode("llm"),
|
|
201
|
+
] = None
|
|
202
|
+
api_key_env: Annotated[
|
|
203
|
+
Optional[str],
|
|
204
|
+
DtoField(),
|
|
205
|
+
BackendField(),
|
|
206
|
+
FrontendField(),
|
|
207
|
+
ExcludeMode("llm"),
|
|
208
|
+
] = None
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
class WorkflowProviderSettings(ModeSlicingMixin, BaseModel):
|
|
212
|
+
default_include_modes: ClassVar[set[str]] = {"dto", "backend", "frontend", "llm"}
|
|
213
|
+
include_unmarked_for_modes: ClassVar[set[str]] = {"dto", "backend", "frontend", "llm"}
|
|
214
|
+
|
|
215
|
+
ocr: Annotated[ProviderEndpointConfig, DtoField(), BackendField(), FrontendField(), LLMField()] = Field(
|
|
216
|
+
default_factory=ProviderEndpointConfig
|
|
217
|
+
)
|
|
218
|
+
parser: Annotated[ProviderEndpointConfig, DtoField(), BackendField(), FrontendField(), LLMField()] = Field(
|
|
219
|
+
default_factory=ProviderEndpointConfig
|
|
220
|
+
)
|
|
221
|
+
embedding: Annotated[EmbeddingProviderConfig, DtoField(), BackendField(), FrontendField(), LLMField()] = Field(
|
|
222
|
+
default_factory=EmbeddingProviderConfig
|
|
223
|
+
)
|
|
224
|
+
|
|
225
|
+
@classmethod
|
|
226
|
+
def from_env(cls) -> "WorkflowProviderSettings":
|
|
227
|
+
def _env(name: str, default: str | None = None) -> str | None:
|
|
228
|
+
value = os.getenv(name)
|
|
229
|
+
return value if value not in {None, ""} else default
|
|
230
|
+
|
|
231
|
+
return cls(
|
|
232
|
+
ocr=ProviderEndpointConfig(
|
|
233
|
+
provider=str(_env("KG_DOC_OCR_PROVIDER", "gemini")),
|
|
234
|
+
model=str(_env("KG_DOC_OCR_MODEL", "gemini-2.5-flash")),
|
|
235
|
+
temperature=float(_env("KG_DOC_OCR_TEMPERATURE", "0.1")),
|
|
236
|
+
base_url=_env("KG_DOC_OCR_BASE_URL"),
|
|
237
|
+
api_key_env=_env("KG_DOC_OCR_API_KEY_ENV"),
|
|
238
|
+
project=_env("KG_DOC_OCR_PROJECT"),
|
|
239
|
+
location=_env("KG_DOC_OCR_LOCATION"),
|
|
240
|
+
max_retries=int(_env("KG_DOC_OCR_MAX_RETRIES", "2")),
|
|
241
|
+
),
|
|
242
|
+
parser=ProviderEndpointConfig(
|
|
243
|
+
provider=str(_env("KG_DOC_PARSER_PROVIDER", "gemini")),
|
|
244
|
+
model=str(_env("KG_DOC_PARSER_MODEL", "gemini-2.5-flash")),
|
|
245
|
+
temperature=float(_env("KG_DOC_PARSER_TEMPERATURE", "0.1")),
|
|
246
|
+
base_url=_env("KG_DOC_PARSER_BASE_URL"),
|
|
247
|
+
api_key_env=_env("KG_DOC_PARSER_API_KEY_ENV"),
|
|
248
|
+
project=_env("KG_DOC_PARSER_PROJECT"),
|
|
249
|
+
location=_env("KG_DOC_PARSER_LOCATION"),
|
|
250
|
+
max_retries=int(_env("KG_DOC_PARSER_MAX_RETRIES", "2")),
|
|
251
|
+
),
|
|
252
|
+
embedding=EmbeddingProviderConfig(
|
|
253
|
+
provider=str(_env("KG_DOC_EMBED_PROVIDER", "fake")),
|
|
254
|
+
model=str(_env("KG_DOC_EMBED_MODEL", "kg-doc-parser-workflow-embedding-v1")),
|
|
255
|
+
dimension=int(_env("KG_DOC_EMBED_DIMENSION", "2")),
|
|
256
|
+
base_url=_env("KG_DOC_EMBED_BASE_URL"),
|
|
257
|
+
api_key_env=_env("KG_DOC_EMBED_API_KEY_ENV"),
|
|
258
|
+
),
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
|
|
262
|
+
def _embedding_vector(text: str, *, dimension: int) -> list[float]:
|
|
263
|
+
checksum = sum(ord(ch) for ch in text or "")
|
|
264
|
+
return [
|
|
265
|
+
float((len(text) + idx + 1) % 97 + 1 + (checksum % 13))
|
|
266
|
+
for idx in range(max(1, dimension))
|
|
267
|
+
]
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
@dataclass
|
|
271
|
+
class _CallableEmbeddingFunction:
|
|
272
|
+
name_value: str
|
|
273
|
+
dimension: int
|
|
274
|
+
provider: str
|
|
275
|
+
|
|
276
|
+
def name(self) -> str:
|
|
277
|
+
return self.name_value
|
|
278
|
+
|
|
279
|
+
def __call__(self, input):
|
|
280
|
+
vectors = []
|
|
281
|
+
for value in input:
|
|
282
|
+
vectors.append(_embedding_vector(str(value or ""), dimension=self.dimension))
|
|
283
|
+
return vectors
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
def build_embedding_function(
|
|
287
|
+
spec: EmbeddingProviderConfig | None = None,
|
|
288
|
+
) -> Callable[[list[str]], list[list[float]]]:
|
|
289
|
+
"""Build a configurable embedding callable.
|
|
290
|
+
|
|
291
|
+
Supported providers currently include fake, OpenAI, Vertex AI, and Ollama.
|
|
292
|
+
The fake provider is deterministic and preferred for unit tests.
|
|
293
|
+
"""
|
|
294
|
+
spec = spec or EmbeddingProviderConfig()
|
|
295
|
+
if spec.provider == "fake":
|
|
296
|
+
return _CallableEmbeddingFunction(
|
|
297
|
+
name_value=spec.model,
|
|
298
|
+
dimension=spec.dimension,
|
|
299
|
+
provider=spec.provider,
|
|
300
|
+
)
|
|
301
|
+
|
|
302
|
+
def _build_langchain_embeddings() -> Any:
|
|
303
|
+
if spec.provider == "openai":
|
|
304
|
+
from langchain_openai import OpenAIEmbeddings
|
|
305
|
+
|
|
306
|
+
kwargs: dict[str, Any] = {"model": spec.model}
|
|
307
|
+
if spec.base_url:
|
|
308
|
+
kwargs["base_url"] = spec.base_url
|
|
309
|
+
if spec.api_key_env and os.getenv(spec.api_key_env):
|
|
310
|
+
kwargs["api_key"] = os.getenv(spec.api_key_env)
|
|
311
|
+
return OpenAIEmbeddings(**kwargs)
|
|
312
|
+
if spec.provider == "vertex":
|
|
313
|
+
from langchain_google_vertexai import VertexAIEmbeddings
|
|
314
|
+
|
|
315
|
+
kwargs = {"model_name": spec.model}
|
|
316
|
+
if spec.project:
|
|
317
|
+
kwargs["project"] = spec.project
|
|
318
|
+
if spec.location:
|
|
319
|
+
kwargs["location"] = spec.location
|
|
320
|
+
return VertexAIEmbeddings(**kwargs)
|
|
321
|
+
if spec.provider == "ollama":
|
|
322
|
+
from langchain_ollama import OllamaEmbeddings
|
|
323
|
+
|
|
324
|
+
kwargs = {"model": spec.model}
|
|
325
|
+
if spec.base_url:
|
|
326
|
+
kwargs["base_url"] = spec.base_url
|
|
327
|
+
return OllamaEmbeddings(**kwargs)
|
|
328
|
+
raise ValueError(f"unsupported embedding provider: {spec.provider}")
|
|
329
|
+
|
|
330
|
+
embeddings = _build_langchain_embeddings()
|
|
331
|
+
|
|
332
|
+
class _LangChainEmbeddingFunction:
|
|
333
|
+
def name(self) -> str:
|
|
334
|
+
return spec.model
|
|
335
|
+
|
|
336
|
+
def __call__(self, input):
|
|
337
|
+
texts = [str(value or "") for value in input]
|
|
338
|
+
if hasattr(embeddings, "embed_documents"):
|
|
339
|
+
return embeddings.embed_documents(texts)
|
|
340
|
+
if hasattr(embeddings, "embed_query"):
|
|
341
|
+
return [embeddings.embed_query(text) for text in texts]
|
|
342
|
+
raise TypeError(f"unsupported embedding backend: {type(embeddings)!r}")
|
|
343
|
+
|
|
344
|
+
return _LangChainEmbeddingFunction()
|
|
345
|
+
|
|
346
|
+
|
|
347
|
+
def build_chat_model(
|
|
348
|
+
spec: ProviderEndpointConfig | None = None,
|
|
349
|
+
*,
|
|
350
|
+
callbacks: list[Any] | None = None,
|
|
351
|
+
):
|
|
352
|
+
"""Build a vendor-specific chat model behind a stable adapter boundary.
|
|
353
|
+
|
|
354
|
+
Supported providers currently include gemini, openai, ollama, vertex, and
|
|
355
|
+
fake. Callers should choose the provider via config and treat the returned
|
|
356
|
+
object as a LangChain-compatible chat model.
|
|
357
|
+
"""
|
|
358
|
+
spec = spec or ProviderEndpointConfig()
|
|
359
|
+
callbacks = callbacks or []
|
|
360
|
+
if spec.provider == "fake":
|
|
361
|
+
return FakeChatModel()
|
|
362
|
+
if spec.provider == "gemini":
|
|
363
|
+
from langchain_google_genai import ChatGoogleGenerativeAI
|
|
364
|
+
|
|
365
|
+
kwargs: dict[str, Any] = {"model": spec.model, "temperature": spec.temperature, "callbacks": callbacks}
|
|
366
|
+
if spec.api_key_env and os.getenv(spec.api_key_env):
|
|
367
|
+
kwargs["google_api_key"] = os.getenv(spec.api_key_env)
|
|
368
|
+
return ChatGoogleGenerativeAI(**kwargs)
|
|
369
|
+
if spec.provider == "openai":
|
|
370
|
+
from langchain_openai import ChatOpenAI
|
|
371
|
+
|
|
372
|
+
kwargs = {"model": spec.model, "temperature": spec.temperature, "callbacks": callbacks}
|
|
373
|
+
if spec.base_url:
|
|
374
|
+
kwargs["base_url"] = spec.base_url
|
|
375
|
+
if spec.api_key_env and os.getenv(spec.api_key_env):
|
|
376
|
+
kwargs["api_key"] = os.getenv(spec.api_key_env)
|
|
377
|
+
return ChatOpenAI(**kwargs)
|
|
378
|
+
if spec.provider == "ollama":
|
|
379
|
+
from langchain_ollama import ChatOllama
|
|
380
|
+
|
|
381
|
+
kwargs = {"model": spec.model, "temperature": spec.temperature, "callbacks": callbacks}
|
|
382
|
+
if spec.base_url:
|
|
383
|
+
kwargs["base_url"] = spec.base_url
|
|
384
|
+
return ChatOllama(**kwargs)
|
|
385
|
+
if spec.provider == "vertex":
|
|
386
|
+
from langchain_google_vertexai import ChatVertexAI
|
|
387
|
+
|
|
388
|
+
kwargs = {"model": spec.model, "temperature": spec.temperature, "callbacks": callbacks}
|
|
389
|
+
if spec.project:
|
|
390
|
+
kwargs["project"] = spec.project
|
|
391
|
+
if spec.location:
|
|
392
|
+
kwargs["location"] = spec.location
|
|
393
|
+
return ChatVertexAI(**kwargs)
|
|
394
|
+
raise ValueError(f"unsupported chat provider: {spec.provider}")
|
|
395
|
+
|
|
396
|
+
|
|
397
|
+
def build_chat_model_for_role(
|
|
398
|
+
role: Literal["ocr", "parser"],
|
|
399
|
+
spec: WorkflowProviderSettings | None = None,
|
|
400
|
+
*,
|
|
401
|
+
callbacks: list[Any] | None = None,
|
|
402
|
+
):
|
|
403
|
+
"""Build the chat model used for either OCR or parsing.
|
|
404
|
+
|
|
405
|
+
Examples:
|
|
406
|
+
- role="ocr" with KG_DOC_OCR_PROVIDER=gemini for image OCR.
|
|
407
|
+
- role="parser" with KG_DOC_PARSER_PROVIDER=openai for recipe extraction.
|
|
408
|
+
- role="parser" with KG_DOC_PARSER_PROVIDER=ollama for local models.
|
|
409
|
+
"""
|
|
410
|
+
settings = spec or WorkflowProviderSettings.from_env()
|
|
411
|
+
chat_spec = settings.ocr if role == "ocr" else settings.parser
|
|
412
|
+
return build_chat_model(chat_spec, callbacks=callbacks)
|