From d3340311084ac9512945fc088ff55e780f676a59 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 27 Sep 2023 21:04:15 -0700 Subject: [PATCH] adding support for completions endpoint in proxy --- litellm/__pycache__/main.cpython-311.pyc | Bin 50033 -> 50660 bytes litellm/main.py | 20 +++++++++++++++++++- litellm/proxy/proxy_cli.py | 3 ++- litellm/proxy/proxy_server.py | 21 ++++++++++++++++++++- pyproject.toml | 2 +- 5 files changed, 42 insertions(+), 4 deletions(-) diff --git a/litellm/__pycache__/main.cpython-311.pyc b/litellm/__pycache__/main.cpython-311.pyc index 64c0ff698f33e7cb95496f33d205e75af1d5befa..3a1ff5970bd702327e271354fbed505d3105bc1e 100644 GIT binary patch delta 947 zcmZWnO-vI(6rOE&`={Mv3;jXB?FC#!Ot1-hz(i=o_)j!m`P;SBY&*^D1_2>WGy#mq zRp|i@UKBy4qyb5Y(WqAqCVOal@}S1!q8>2u;G0EK(3hDn^S{ zuLpr`Y*fnG-;4a)xsV%NAIu+@8HA>g#Ede7h)n6<=K5O3?4x#xi_jgGNe=LxWgk4( z4$rNBTnyGoyUL`ryS#>QE5I*#%cO+1=&ZzVqkYqCf&r;w8@FfLF~*Lv7m>zowSAU!jh~s z{Y}!=Fd!=m7Or8UNDX3-XeXCwW39B~5UfOww9)l33Y~{3b^w$-L{>96t;*dkD!ple z8nJYRCZKQ^1ogYVGT)vJ)aJCcK=aHX0E|i3Cx3M=cBi=*Y$%5I&A5$FW3j3>moTd8 zO$UlZXCg-U!0d4&+&nvMgpaO>MmSysS~Ma}vt35yP~K@ojudNzVpNE^t=i742eVcx7j!MfOH=DmE-VgE*!ossghIw5mJ)dK^>%hD#+{*hn8^QYeQb5ly@tpOb`unAjm%c1q z8~s0)eu2Sy`|=2r)n6`4+ytl}3iU5;vO$`fJfvYXWrRZdRa5x`hEdDD(6#@flcJ5yL&fwFrxPutndA_lUdXdy@$ t$mClrIhn;JMXi(d_sMci0}22ULvi!wuzf7-{H%=H9~j_-z~r2x6#&jbVebF{ diff --git a/litellm/main.py b/litellm/main.py index 20e554c7687..160f786f5f5 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1431,7 +1431,25 @@ def text_completion(*args, **kwargs): messages = [{"role": "system", "content": kwargs["prompt"]}] kwargs["messages"] = messages kwargs.pop("prompt") - return completion(*args, **kwargs) + response = completion(*args, **kwargs) # assume the response is the openai response object + response_2 = { + "id": response["id"], + "object": "text_completion", + "created": response["created"], + "model": response["model"], + "choices": [ + { + "text": response["choices"][0]["message"]["content"], + "index": response["choices"][0]["index"], + "logprobs": None, + "finish_reason": response["choices"][0]["finish_reason"] + } + ], + "usage": response["usage"] + } + return response_2 + else: + raise ValueError("please pass prompt into the `text_completion` endpoint - `text_completion(model, prompt='hello world')`") ##### Moderation ####################### def moderation(input: str, api_key: Optional[str]=None): diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 75bad1b6fb2..82856e45898 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -7,7 +7,8 @@ load_dotenv() @click.option('--api_base', default=None, help='API base URL.') @click.option('--model', required=True, help='The model name to pass to litellm expects') def run_server(port, api_base, model): - from .proxy_server import app, initialize + # from .proxy_server import app, initialize + from proxy_server import app, initialize initialize(model, api_base) try: import uvicorn diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4535775d216..a1934242705 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1,4 +1,10 @@ +import sys, os +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + import litellm +print(litellm.__file__) from fastapi import FastAPI, Request from fastapi.responses import StreamingResponse import json @@ -25,8 +31,21 @@ def model_list(): object="list", ) -@app.post("/chat/completions") +@app.post("/completions") async def completion(request: Request): + data = await request.json() + if (user_model is None): + raise ValueError("Proxy model needs to be set") + data["model"] = user_model + if user_api_base: + data["api_base"] = user_api_base + response = litellm.text_completion(**data) + if 'stream' in data and data['stream'] == True: # use generate_responses to stream responses + return StreamingResponse(data_generator(response), media_type='text/event-stream') + return response + +@app.post("/chat/completions") +async def chat_completion(request: Request): data = await request.json() if (user_model is None): raise ValueError("Proxy model needs to be set") diff --git a/pyproject.toml b/pyproject.toml index 3ffbfd1a760..71d5cdf42b6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "0.1.789" +version = "0.1.790" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT License"