Refactor OpenAIModelComponent configuration fields

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-07 21:00:03 -03:00
commit f21ddc2306

View file

@ -11,21 +11,19 @@ class OpenAIModelComponent(CustomComponent):
def build_config(self): def build_config(self):
return { return {
"inputs": {"display_name": "Input"},
"max_tokens": { "max_tokens": {
"display_name": "Max Tokens", "display_name": "Max Tokens",
"field_type": "int",
"advanced": False, "advanced": False,
"required": False, "required": False,
}, },
"model_kwargs": { "model_kwargs": {
"display_name": "Model Kwargs", "display_name": "Model Kwargs",
"field_type": "NestedDict",
"advanced": True, "advanced": True,
"required": False, "required": False,
}, },
"model_name": { "model_name": {
"display_name": "Model Name", "display_name": "Model Name",
"field_type": "str",
"advanced": False, "advanced": False,
"required": False, "required": False,
"options": [ "options": [
@ -39,7 +37,6 @@ class OpenAIModelComponent(CustomComponent):
}, },
"openai_api_base": { "openai_api_base": {
"display_name": "OpenAI API Base", "display_name": "OpenAI API Base",
"field_type": "str",
"advanced": False, "advanced": False,
"required": False, "required": False,
"info": ( "info": (
@ -49,14 +46,12 @@ class OpenAIModelComponent(CustomComponent):
}, },
"openai_api_key": { "openai_api_key": {
"display_name": "OpenAI API Key", "display_name": "OpenAI API Key",
"field_type": "str",
"advanced": False, "advanced": False,
"required": False, "required": False,
"password": True, "password": True,
}, },
"temperature": { "temperature": {
"display_name": "Temperature", "display_name": "Temperature",
"field_type": "float",
"advanced": False, "advanced": False,
"required": False, "required": False,
"value": 0.7, "value": 0.7,
@ -85,4 +80,6 @@ class OpenAIModelComponent(CustomComponent):
) )
message = model.invoke(inputs) message = model.invoke(inputs)
return message.content if hasattr(message, "content") else message result = message.content if hasattr(message, "content") else message
self.status = result
return result