jeff3071 commented on code in PR #71437:
URL: https://github.com/apache/airflow/pull/71437#discussion_r3766395515


##########
providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py:
##########
@@ -160,30 +181,57 @@ def get_conn(self) -> Model:
 
         provider_kwargs = self._get_provider_kwargs(api_key, base_url, extra)
         if provider_kwargs:
-            _kwargs = provider_kwargs  # capture for closure
             self.log.info(
                 "Using explicit credentials for provider with model '%s': %s",
                 model_name,
                 list(provider_kwargs),
             )
-
-            def _provider_factory(pname: str) -> Any:
-                try:
-                    return infer_provider_class(pname)(**_kwargs)
-                except TypeError:
-                    self.log.warning(
-                        "Provider '%s' rejected kwargs %s; falling back to 
env-var auth",
-                        pname,
-                        list(_kwargs),
-                    )
-                    return infer_provider(pname)
-
-            self._model = infer_model(model_name, 
provider_factory=_provider_factory)
+            self._model = infer_model(
+                model_name,
+                
provider_factory=self._create_provider_factory(provider_kwargs),
+            )
             return self._model
 
         self._model = infer_model(model_name)
         return self._model
 
+    def get_embedder(self) -> Embedder:
+        """Return a pydantic-ai ``Embedder`` using this connection's 
credentials."""
+        if self._embedder is not None:
+            return self._embedder
+
+        conn = self.get_connection(self.llm_conn_id) if self._conn is None 
else self._conn
+        extra: dict[str, Any] = (
+            conn.extra_dejson if self._conn_extra_dejson is None else 
self._conn_extra_dejson
+        )
+
+        embed_model_name: str = self.embed_model_id or 
extra.get("embed_model", "")
+        if not embed_model_name:
+            raise ValueError(
+                "No embedding model specified. Set embed_model_id on the hook 
or the embed_model field "
+                "on the connection."
+            )
+
+        api_key: str | None = conn.password or None
+        base_url: str | None = conn.host or None
+
+        provider_kwargs = self._get_provider_kwargs(api_key, base_url, extra)
+        if provider_kwargs:
+            self.log.info(
+                "Using explicit credentials for provider with embedding model 
'%s': %s",
+                embed_model_name,
+                list(provider_kwargs),
+            )
+            embedding_model = infer_embedding_model(
+                embed_model_name,
+                
provider_factory=self._create_provider_factory(provider_kwargs),
+            )
+        else:
+            embedding_model = infer_embedding_model(embed_model_name)
+
+        self._embedder = Embedder(embedding_model)

Review Comment:
   I passed instrument to `Embedder`.
   Thank for pointing this out!



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to