fix: remove double await in ChatOllamaComponent (#7992)
* Fix double await in ChatOllamaComponent * ChatOllamaComponent: fix respect given URL self.base_url * ChatOllamaComponent: implement support for "load_from_db" variables * ChatOllamaComponent: fix RUF005 problem * [autofix.ci] apply automated fixes * conditionally await if variable is a coroutine * Remove outdated comment regarding coroutine handling in ChatOllamaComponent * ChatOllamaComponent: await coroutine result for JSON data retrieval --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Gabriel Luiz Freitas Almeida <gabriel@langflow.org>
This commit is contained in:
parent
c2e78180f1
commit
9f7b4be36d
1 changed files with 27 additions and 15 deletions
|
|
@ -1,3 +1,4 @@
|
||||||
|
import asyncio
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from urllib.parse import urljoin
|
from urllib.parse import urljoin
|
||||||
|
|
||||||
|
|
@ -219,20 +220,27 @@ class ChatOllamaComponent(LCModelComponent):
|
||||||
build_config["mirostat_eta"]["value"] = 0.1
|
build_config["mirostat_eta"]["value"] = 0.1
|
||||||
build_config["mirostat_tau"]["value"] = 5
|
build_config["mirostat_tau"]["value"] = 5
|
||||||
|
|
||||||
if field_name in {"base_url", "model_name"} and not await self.is_valid_ollama_url(
|
if field_name in {"base_url", "model_name"}:
|
||||||
build_config["base_url"].get("value", "")
|
if build_config["base_url"].get("load_from_db", False):
|
||||||
):
|
base_url_value = await self.get_variables(build_config["base_url"].get("value", ""), "base_url")
|
||||||
# Check if any URL in the list is valid
|
|
||||||
valid_url = ""
|
|
||||||
for url in URL_LIST:
|
|
||||||
if await self.is_valid_ollama_url(url):
|
|
||||||
valid_url = url
|
|
||||||
break
|
|
||||||
if valid_url != "":
|
|
||||||
build_config["base_url"]["value"] = valid_url
|
|
||||||
else:
|
else:
|
||||||
msg = "No valid Ollama URL found."
|
base_url_value = build_config["base_url"].get("value", "")
|
||||||
raise ValueError(msg)
|
|
||||||
|
if not await self.is_valid_ollama_url(base_url_value):
|
||||||
|
# Check if any URL in the list is valid
|
||||||
|
valid_url = ""
|
||||||
|
check_urls = URL_LIST
|
||||||
|
if self.base_url:
|
||||||
|
check_urls = [self.base_url, *URL_LIST]
|
||||||
|
for url in check_urls:
|
||||||
|
if await self.is_valid_ollama_url(url):
|
||||||
|
valid_url = url
|
||||||
|
break
|
||||||
|
if valid_url != "":
|
||||||
|
build_config["base_url"]["value"] = valid_url
|
||||||
|
else:
|
||||||
|
msg = "No valid Ollama URL found."
|
||||||
|
raise ValueError(msg)
|
||||||
if field_name in {"model_name", "base_url", "tool_model_enabled"}:
|
if field_name in {"model_name", "base_url", "tool_model_enabled"}:
|
||||||
if await self.is_valid_ollama_url(self.base_url):
|
if await self.is_valid_ollama_url(self.base_url):
|
||||||
tool_model_enabled = build_config["tool_model_enabled"].get("value", False) or self.tool_model_enabled
|
tool_model_enabled = build_config["tool_model_enabled"].get("value", False) or self.tool_model_enabled
|
||||||
|
|
@ -286,7 +294,9 @@ class ChatOllamaComponent(LCModelComponent):
|
||||||
# Fetch available models
|
# Fetch available models
|
||||||
tags_response = await client.get(tags_url)
|
tags_response = await client.get(tags_url)
|
||||||
tags_response.raise_for_status()
|
tags_response.raise_for_status()
|
||||||
models = await tags_response.json()
|
models = tags_response.json()
|
||||||
|
if asyncio.iscoroutine(models):
|
||||||
|
models = await models
|
||||||
logger.debug(f"Available models: {models}")
|
logger.debug(f"Available models: {models}")
|
||||||
|
|
||||||
# Filter models that are NOT embedding models
|
# Filter models that are NOT embedding models
|
||||||
|
|
@ -298,7 +308,9 @@ class ChatOllamaComponent(LCModelComponent):
|
||||||
payload = {"model": model_name}
|
payload = {"model": model_name}
|
||||||
show_response = await client.post(show_url, json=payload)
|
show_response = await client.post(show_url, json=payload)
|
||||||
show_response.raise_for_status()
|
show_response.raise_for_status()
|
||||||
json_data = await show_response.json()
|
json_data = show_response.json()
|
||||||
|
if asyncio.iscoroutine(json_data):
|
||||||
|
json_data = await json_data
|
||||||
capabilities = json_data.get(self.JSON_CAPABILITIES_KEY, [])
|
capabilities = json_data.get(self.JSON_CAPABILITIES_KEY, [])
|
||||||
logger.debug(f"Model: {model_name}, Capabilities: {capabilities}")
|
logger.debug(f"Model: {model_name}, Capabilities: {capabilities}")
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue