From 7e34736a3836fe53f3ee07bbf8676e199e0c221d Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 6 Oct 2023 15:01:50 -0700 Subject: [PATCH] fix(add-custom-success-callback-for-streaming): add custom success callback for streaming --- litellm/__pycache__/main.cpython-311.pyc | Bin 49666 -> 49752 bytes litellm/__pycache__/utils.cpython-311.pyc | Bin 153580 -> 153591 bytes litellm/llms/prompt_templates/factory.py | 26 +++++++++++++- litellm/main.py | 3 +- litellm/proxy/proxy_cli.py | 12 ++++--- litellm/proxy/proxy_server.py | 12 ++++--- litellm/tests/test_streaming.py | 40 ++++++++++++++++++++-- litellm/utils.py | 16 ++++----- 8 files changed, 89 insertions(+), 20 deletions(-) diff --git a/litellm/__pycache__/main.cpython-311.pyc b/litellm/__pycache__/main.cpython-311.pyc index caba72a447c48a7dfb2ce9733b0220afb4f4ed9a..b4ea2793100d89325c03f6650d005cefdfe2f247 100644 GIT binary patch delta 908 zcmZ`$Ur1A76#ve?-MKcenzrdQak@EYr8YI4f3oh~24+#{&|Z2dnDh|kyjrN-62yY) zq2LxqT0IClH`C&?pxzeLW9TMiMOZ{aJ;V&mB8a~4TDX@E_nzPR{m$a?Vq z?fjoFJ;|3bracMNyERG69MhfG%343_O;L|htoKaXAt>FZyb4)U2MkhYPq6DuPiMVV zSr4tHX{8qWEm9-(O0A%*P5CtjP+G!m+3)oHd|L4>H%;_e&^s3Oj(NRcU>pP^#^~(-&NOKt6|;TQuB9a-L#TQ#qFUU3MP7}hgc+x#uT$4 zf())8iM<%8yP6nBB6^#k9s|t>qX&yf@Ff<*q=IPfEae`!^PTVK+@p}&H@s)drhh|*zga9V-N2L#d!~Io-0R51N#3LWwMqQ@-XXzKUBzYLsYHk=So#%{7bD`#&9#I)64K%i;~MQo>O%FFp2xE ztMtMAA(?Dv1WTE=>cLlYa;ZHen)$c~NmGM7~s^uF>j^bm25wt5-(rEXZ8dE|o zL8sN6Nbx7IA!NO8oGDjrmeO6-Znv5~Sx~~6F5c*?%T~eyZ-fQA5_X~I_ND6dFc309 zG|=G&iA+#BtCi@yW|#i`U7|Tmp0da3!MP%OYfgeR>YK|S(u6%r6O(#kj-@8?5 zj3K!jjylOCCK{uT_~_)jnopxyjhTtB5u;{$;%+jDi6$|@5d&tkajQ-v?e1j%+3i1m zUAInEojP^u)Tw*F-5=1iKft`pY!1}nZ_9>6CGVhl!VGDozdCrqPZh|3%z0$P`&*vOFXqAswQaXqpX4gw zpZP(;fPx33IvpF}DpW(Big}?X(9ove6so>7d)PnfPjLGcZVur4a`RsV~TmX9JWT{^$; zm}^lfSEELDMvYXD&yRpn>XrHXgB^fcU8LIEHBC}qd1AtpZ1nQ^Jbh`p?qv9^v=T`= zuaofMycAeEOgq4@OtP8_P_Cos5vO)R^S>3*MkUMTcxDjrtV;0ST z&)rRnPD#e^(I~s8<}ThGT@G?+m+l|XE~&Q`N4oDXJ_*pHo?NmRG1)d+{s}b`+>hCg z`p4fVBdr9d5j;j)LxWOZ$*vRUu6y+IK>*X;`L$&RRBd-8F72TK6X_3eUe+7d98dJ@vKb>;&HoFOZ>I1ph)H zCN!$|numw2mi1RD;*IN>y*u8p7 zJAg?Yw5?>}GU9Ki^;99{k!>xFOkrEZf{77S-B?wHv9~U>RVZvJu}m&UmEx*anAKJz zHj&jJ+Z&|iMAZ`5)uY=c$PGw7>+ap=2$D>oKKoIWx^CAJoceFOs^JrL_M3;H&TZH| zNiwY<9Y30Ol{#~8f#0L_tOK=uZ>ebw(IWgS)ULfNF>`DFkPcxkVmA^rAy|Mtt?t@4 zalT*^hqd$zG!wOwU<^$`lX(WzH`ZCJlzN5PTs4k*rw1Cu`Y@8qNyulXQ61Z13xDX@ zXGq=Mk&)!U=U>q3mugmfIuaAa3V+Fn*Ex0k{-rTuSS(4nAXJ3MP3ocjV@hEY@{+R|=sm%u; zUHd4KmYDySwN+$Im=p;{HA&PLPq4E=si?8ptqnG2t8+4ugllQ|w**_!l~DFHk*x&3 zA=s!|59P=jk_X(|59OON=lSgs5OwIpe8c?pB(Z^D6Tx!`KfYB6*4mtk3#B?}{DAw7QyU;8hv*ovi%(*kiWM=mBj;XJli!OMVCLb_8cYc0(eo6lrt=JQsS6Y-m zFTbD3i8-)nL1E#{lEQ-AJl{x%`t!NSXc3vvdHH3@F{XDzyCnEoo;AaAn02hd&n6OI(v`}`1;3$G8j5c9j>4Nh7dHu24r3QUrSxXUP z9}$QR2%-(>7jY2(D#0;=iv+YAtdl@w&1Ix4;p{juY_zTuNO>fO)8nVquc%>;daS$1 zF)AxD(<<3}B>1lcBGIaRnRb$xrwHCh@PzeU6?9WP$hgyD;GqlS!2sna`+$Iw5=)gh zem@f9KL|`}?DZl!0K@sYyZZY4P%x`Ue<(HTk--dVXwTq~Fd_#M#1KTQQ+rMt|0vgy9ij$Qiy($;39z;#XO--r5+m&tBaAtcnc_AbJ1ku9Ev&p5wN8YKW|^5 zC*KIxyEpY-4(;oncvT?GgDCz^AUt9!BoQ~w62`#<8(=7Z(gazi+r;`i!5x0e1Sych zZ>$n1N|cS@E&(9V6QE#z z(|x9};lvz8D@o!FW{4;lP4rpfQq^P>J&#|7t-(e2y@mKA2o!=J2*#7bmqd*t>LyWR z`B^gz8%>=bGoTqul(z*F4bcn4Jx}24p}gHC9}y7@0PISO(>FV)}S`O)wSysll320cKMb~6T4fXDdiIB0}8ekTr&n2)ats7$Fv z=g{C0?MOT@NX3?&W80raOz$s6$7ZY!bxnBitkb!&`H=+3Fi-ZemJg=4f{}w)hpq`l z6Q5StbieWd@??zaBDJ-eZyBKL(bYnvD4n&@Fjx{SK{rqh0PMODKfyfvEE#u z*w3L+kq)h-HU=L(CcC4`>TIm2P@GOSmITDkXAvc~e=Lc9pu<-X_9r66RZNGA^(O;6 zPLC%Djv-jW9?FAT{GWzE)PvPOHw5NHMUyWC>#%S-ugW@ItFFti7I^m%cxOO01ccw# z5k6DYB~@w1QXs%!xYv7kkY=`EN=%&}p*uOSBy?`7bTLK3hl{C!^RncNqbJTAC3jkk z$mtvdK=YZfeG=uQKKb%zO72Bv(+jr ze`STzK+SbxMk3qA6P7^*?BE-gfdii7ca}jraK5n$6z@I$62Dso-YitZ_osoLdg zusU5fSEG}K)4I>_-`B!0EGDOGA$!JfviLs|S_I8|<?KSwqYp=8VUVDW#gV&l^qUb(a&{3xWArXJx zD=3)E25WzM3cl8x#^Zwrnk;6f=kIQWc$lVrvJskO$mQj2Fcj=sOB-yEqCF%RByB%%c^l5GNVM#74gSMOfHJYahM{1Enw{NqvmYSR^(L|5%~+ zJqNQCMLFp))i|xoYU*p8)%f}`iiW%SplvV{U0J*h(x6b=xDED5P{_x<3?nSvWXHh1 z#lU$Cs8{JNY;D92@S6XQf?B9O^D=y<_ulM<+_4=pCJ0?Wq3gW|Z;_~FEW(9EUmGaWUi5jQccVU!Z6`U=)NP>04&HeN#^voIS~QV+h@$g%b1>Se+P0N` zyQi`P#NR(~?od~Xy7*u-a@#cE{Y%2}8eE?hFwr;-5#Hj;jezjGGGdy68q z2OimJSCoc_j)O=Han;B0h;xt(_xQwfFdKTcC(prB6a1*%_!#~mnZH2~aL@0MBcN^n zEBqM}D1-6kxxf5XG*z-`Bp|kPI&Zs%m!DqSeGPv0#|s(1rL8w_s0(KY`QcuyDfhK&y)eoTYF;!;Hse7Wk=}Mm zB?i1rP8y_kT*E?H+6c53qqGRX$#eas419F>OZbtgiN7Bpy#i&L5-456Z%bNUkaWO^ z?Cvm0$)K=`Y(PJSi#|r*Ew+FpD#^7n{?tHeBpQBupmb5c{978>i31-0{+AHCH(CO6 z=}S^Zaq%3Bsjt>h28-RWk)ac$ISj!gH{$0)=6&l3^S-KDmo{4)m}5mvB@T$F%UV~Y zbXktJqlN`M11gnewnn>)zDv0r_}TE5_HXf0hhLKTU?)Pi6T|5VY-n^@>m5zD8W$5M z)6Vw|l?tF+>mDlA=#8_;sd_lC-54ny@sAT;%Ir1T(boxf5RgaAt{s~oT?aU*y*^P2mC%#7 zCreNInMBM_k;Dp~pC?TpC+;y(UqlyT@;V}{oggvMpa`uMw1_9j!y@fyo-{`SJr69v z?eEsI3#1a5oK8X0i+w67K|-7!QEvz=a>;}9geGV+nNS!|P{Jztdxg@R^fZ+77%Ls3 zI}u?LJx3%>Y{4<_4%;=#bfXG;s$incR}=1(ISLm<9@E+*=IGD068 znMOe0IT(E*@E(7W$UXG<2F*xsnc}5Ryq1Wblr1Kn=vrO$s8i? z;tz@Q4uMBA6-!6t%sd(JL;5e{ZNcw;c(f!?yE|8UUJA>Vz}%&K2#d5W-A delta 6349 zcmbtYYkX8ivOm=`XOhfh67nJ;A!I@lU=jj`@DfB4@VXWu)%)e%`S71Q z)m7Ei)z#H?&dn`hSGI&%|J`aeEBMP>nc|AwW*t=+`hJL_2*nE`XB#B?>I$C)1DpAo z?`qMP#?a;Lh52cU_dcz1Qi^pv=>>kXZDQM4$kc`mu^Lmo6BI3KNWaJlcSaS3M|umi zw;I!pXyAf&yfId@mBswnz$o9mvSkeVYo>~m+E*37W`8hh`7@GODQGii*tKJozp;+8 zW?xefmWOt-k-qjRSDCh<>bDUNKwc*eba%Fb_F(n>!^?*&d*deNPByb6p$s36n9b!w z)g!U_y~;DyqvrLS?VE=S;S=B38T**&Dz1WeYpP>e zpE6LRI+Y*c45n>!#QXL-_5ysPJ?DG`@pI=r>UXFa?@OAsD>V6gGSWftA%frJY;AS5 z)bdNjx$e6@I~HJ!FRZ!Jh^o_DZKiK%z^rw&_AXWc6=FTkI+vq{J3S2!?rL>&Z(~cn zUFAj`|uO}*1f8~>1V`wP_w;agF;{aD=Pua+WA+@r@M&1hSpPulwWnWw{aI= zMT~l4gg3O+*JJD*bW(VoQ-L z{3#;==d@c}W+lpDxg_a=RFNLf)#h%^wFHOfYTLIC%o#-!8}9OywEQ8nT(#U?(}c%G zhCymcD_$*p+rYx05C2<$?4bco>4bDg>c{{}Z`-ta+Xkda>%mw_!MkO4ZL|M8=t0l2 z>=FD?iVY_SE=ev{S_`b0MuXO*tE*{W+G*f%ZQMIQTUd>xJ@Nl#?QyauP09q5=}h&- zAK__r)igSr9j#98oa5m#2^Wy`5`tGyCWV42-mB)PlI>*9PJiLBNh*q>uUCJAB}Ix>-mge`2?52^(3Hj5Cs*1{cxU^t@L_tOAgGEz|dblg4mD1{xl8NOd#rXxn z5n8uDs=TbyQBYPEsIqKgL4H+GzD4``=sD9t)Z&MGK9!|jLdIL#j8jvRWW{=$_Bo8oK#j8*Lq-*_#w0olA=dZDQ!xxa&^piDZ;?46 z73#Em&yL0$i=Q1A7vz!|{y9xsr|mnNi??a_*}>{J$k^typR-z%cG8%vZ45=LxD5Xh zVm512{xZ7cJ0kyw;2#7sJB|`X5#uKbXwCc#!S50L(X?$-E2^uCruKwoop$^$_JtHq zeu_YDSvYN3k4Ovg&l1qy@#6%vSNsHl%(K%-+2i>8#BkEWK0wOP+#Y|3t7UE@ceh|w zMvlq#7$#dS-$R0bAdsmU7tGE7B<5a%ecFqc62|OD(jR@--Ek*VDSi2e1eCK_wd6_s zmKe7P%-Uy{O4U$|-6y`mpHGW|aINEJ#c+lUZXmc!4jCI;8ohikxqO)55J40TM-%iR zNYD)3drf_b%+L;Z_l@v}*R*-O?m1=d`g#}F{?;9z9*(2_D37bf%jY}19B(}bUJ$N; zl>|~fqKMD3Au7a*ypvrXj}wfq)fC~16KN21Cohm5J~=oYQ}s2nK*^eQgqG087Q_WN3p*u{jyu zjqG|NtTw#~eM5Wic~iuvCP)^?lAv8!Q{aNJCD%IwTkk|8Pxp=#Tl->ApG}4V8HN`Q zVJAPD)B?YBg%>^<2ANWk)!t8Z zhvTw`_lGLzFP`rYgA&t7w1A+HU=yy3ZzlLs9O@6_7N&BAwm@!ZFn63Jafv|AS5B&f z=+RD)Nzju>d@_y7-07f!9DMX!n%(sdPg_k5rWPMY0y24WiIN$VOQHu9@mv~M@7YZc zBWdh?qIxole?X6W5bV)+< zvK7MOt|^4SEbU}&-I4)e#*$7YEUx#BSfzKrRShTN2A+-^*r|j>^}ZNuI~SXIGB)#k zO2+w&tn+Ev=Tq$$2jq5Il(?*GiYc!5b%IWUZ@X_q=m~aA%kfoaWp7M{IyI4B%{_TK6k!tS9t`I&&!5hQ0R!YV711)sAYd~l3*t)@AX!jKXlSW8mJJyud1P#@oQ zG&5+E`ik6fFf=QVtV|$~4s0VzmR{NK2e;1559%+D1G^e0tMhYY!bxu2Hg;$NB=-s$ z>yS&*yC;C(1|8y54Xp1)_Zx2@#b%-wiKVqz3YUuAwJ-+`iYyn{E%J_+L2(L)3+BQG zan=Q+V7-W_gCtldGV7oQ-Vn>`zzxg9_jRxmR*Dt%;0m+{t5-HaRe;P-CV4Y< zi9a>KJFrV%(FlvcBJ1uGHxVHC`^mAkWGKQQHLB33TQ8Ee&38 ztHa}UdfPlaj#hY3Y-oakkR(nt!RW$)Wbt2;Uxp`mq4lPf#}h0k;XX*&dj+y!Eiueu zax-2Uc{;{OO^mmYWHbq>fukGH7t0ny{7|WE3@FPL$;PD}{qh(QZ1@YjGF!=tRX@EL zzBE`y;DaCL*{z(3kC#9)jMu+d0`pYJ6OQGO0Zsa|%i(F(=K{%vdt5c#<#o(;Hn+KW zU>PD+gTW>m<`?v78uVo*Ik76*R>1Vzw0>p<#IR_Uq|T!=cEwJ6`IBNZr7l`Ztoy6YsBq6rAVNHE@{bd3!A;S)Mqs zR&u{s3m>V--TD?j=IO`Zf-x+$ov<_7BzwkUvE~rO*qvlT zCbVn@WrIh}AYVhPk)7XD^tee}JOm@M*AXo{%?(7YN4VS%=d6GA75Z%$&$kf&4T5Mf z{xHP#l{+ajF^j^Q#UCOc<$*Md*29p8&TTskG2j#X55venL35#uU-dM*T&+Bq;BxO} zx6ntV9)bSQEeemoB)F+BKLWEXrf+G2IPpUfi_@(a;ALk05*@(}{+j$z^o^guaTrbM zif`aOMXY9CEgw$;^2p{BMJLW~@>^=$Ep?6c_*Um_;7s3s5xx(#%te7lB^0G9#qQ7H zFgV1LFCYzab*#-Uh?6~%jJ=ZGL``x=(2IN8RV@J3$> zNbnl~eh-ivA+L+?MN>Cc&uZ~nH`q+nxbv-IZ#UMC@AbcS!_W|DSdqY-CfP;8e=KC> zM!X-68re!*kxOMufL3N=4+C&f7Rs{mu|1UGSF5?=co=&fD)d$}y9l7@<>73r3E4kH zGgmeRRJwuQNJoN<+xJoxNz{@<72>5BHUtgtiD4%UlfR;YH*vrp+VkE*H$i(?K7E?X zE-RaCx7;o=6iT@_PBKJqDtrKfUv0zBjT~Q8#HaDBFJ86Z#j`UiG>H?*EFP}tmy_A% zkW~3;PsZd;jH=(<+U9k%xaT?XHA0@bCb2Vv6+@T)Lk4R!m>wYWH}w`fi-g?Z2Ff(s zho0c^%p*6bv~m*mTV{LQEqD`f{C>tK>02{d65#BIv)BLw9MQXmuy;ezsZFC;A{Nw- zMzKo(TlLNNvM7eL?itIT4UwtI_tCDmiOK>tVT8PNWR;N3fh92M8zixZK=utX+->CW zO!B5gKTyCXGhiaF827MCFDPc^Fg6PXaDyn=<>g3Jk^-(I0Cu%Jz@c#|#HB)HaXGIM zhfCPx!I>!RH`Tgj(;!14TLGEGgJ^T)UEn1BMWU#bO@VPjlrkF>={rl=Nq{J^^Z_>J z;lU*PFK=zK>6EP|b(#DtG9vp%*;vXxQ+Ai~b>#q!9wbO7kazOy6b*y0O=8nsvdy&7 z@-hi>2;}F|nMCa;9r{wqcM#Br5>DR^0ta0t={NMajvS%)Q2Ev;-(_SY${!)#4e{9| zHn5kxMh_C>J%S5*?=tqTn)A33@KgFv<2A{@e)s5tupLv`ODuXK18b*p7yev-ql#4- F{s~wwYoq`G diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index d2454848872..3653c7256ce 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -1,6 +1,7 @@ import requests, traceback import json from jinja2 import Template, exceptions, Environment, meta +from typing import Optional def default_pt(messages): return " ".join(message["content"] for message in messages) @@ -26,6 +27,25 @@ def llama_2_chat_pt(messages): ) return prompt +def ollama_pt(messages): # https://github.com/jmorganca/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template + prompt = custom_prompt( + role_dict={ + "system": { + "pre_message": "### System:\n", + "post_message": "\n" + }, + "user": { + "pre_message": "### User:\n", + "post_message": "\n", + }, + "assistant": { + "pre_message": "### Response:\n", + "post_message": "\n", + } + }, + final_prompt_value="### Response:" + ) + def mistral_instruct_pt(messages): prompt = custom_prompt( initial_prompt_value="", @@ -190,9 +210,13 @@ def custom_prompt(role_dict: dict, messages: list, initial_prompt_value: str="", prompt += final_prompt_value return prompt -def prompt_factory(model: str, messages: list): +def prompt_factory(model: str, messages: list, custom_llm_provider: Optional[str]=None): original_model_name = model model = model.lower() + + if custom_llm_provider == "ollama": + return ollama_pt(messages=messages) + try: if "meta-llama/llama-2" in model: if "chat" in model: diff --git a/litellm/main.py b/litellm/main.py index 2589c85fb32..45537eb494a 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -961,7 +961,7 @@ def completion( messages=messages ) else: - prompt = prompt_factory(model=model, messages=messages) + prompt = prompt_factory(model=model, messages=messages, custom_llm_provider=custom_llm_provider) ## LOGGING logging.pre_call( @@ -1410,6 +1410,7 @@ def text_completion(*args, **kwargs): kwargs["messages"] = messages kwargs.pop("prompt") response = completion(*args, **kwargs) # assume the response is the openai response object + print(f"response: {response}") formatted_response_obj = { "id": response["id"], "object": "text_completion", diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 28f252fe44f..9867811bb3c 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -61,7 +61,7 @@ def open_config(): @click.option('--max_tokens', default=None, help='Set max tokens for the model') @click.option('--telemetry', default=True, type=bool, help='Helps us know if people are using this feature. Turn this off by doing `--telemetry False`') @click.option('--config', is_flag=True, help='Create and open .env file from .env.template') -@click.option('--test', default=None, help='proxy chat completions url to make a test request to') +@click.option('--test', flag_value=True, help='proxy chat completions url to make a test request to') @click.option('--local', is_flag=True, default=False, help='for local debugging') def run_server(port, api_base, model, deploy, debug, temperature, max_tokens, telemetry, config, test, local): if config: @@ -82,10 +82,14 @@ def run_server(port, api_base, model, deploy, debug, temperature, max_tokens, te print(f"\033[32mLiteLLM: Test your URL using the following: \"litellm --test {url}\"\033[0m") return - if test != None: + if test != False: click.echo('LiteLLM: Making a test ChatCompletions request to your proxy') import openai - openai.api_base = test + if test == True: # flag value set + api_base = "http://0.0.0.0:8000" + else: + api_base = test + openai.api_base = api_base openai.api_key = "temp-key" print(openai.api_base) @@ -107,7 +111,7 @@ def run_server(port, api_base, model, deploy, debug, temperature, max_tokens, te except: raise ImportError("Uvicorn needs to be imported. Run - `pip install uvicorn`") print(f"\033[32mLiteLLM: Deployed Proxy Locally\033[0m\n") - print(f"\033[32mLiteLLM: Test your URL using the following: \"litellm --test http://0.0.0.0:{port}\" [In a new terminal tab]\033[0m\n") + print(f"\033[32mLiteLLM: Test your local endpoint with: \"litellm --test\" [In a new terminal tab]\033[0m\n") print(f"\033[32mLiteLLM: Deploy your proxy using the following: \"litellm --model claude-instant-1 --deploy\" Get an https://api.litellm.ai/chat/completions endpoint \033[0m\n") uvicorn.run(app, host='0.0.0.0', port=port) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 356ace1a868..36b96c754e1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -18,10 +18,12 @@ print() import litellm from fastapi import FastAPI, Request +from fastapi.routing import APIRouter from fastapi.responses import StreamingResponse import json app = FastAPI() +router = APIRouter() user_api_base = None user_model = None @@ -109,14 +111,14 @@ def data_generator(response): yield f"data: {json.dumps(chunk)}\n\n" #### API ENDPOINTS #### -@app.get("/models") # if project requires model list +@router.get("/models") # if project requires model list def model_list(): return dict( data=[{"id": user_model, "object": "model", "created": 1677610602, "owned_by": "openai"}], object="list", ) -@app.post("/{version}/completions") +@router.post("/completions") async def completion(request: Request): data = await request.json() print_verbose(f"data passed in: {data}") @@ -149,7 +151,7 @@ async def completion(request: Request): return StreamingResponse(data_generator(response), media_type='text/event-stream') return response -@app.post("/chat/completions") +@router.post("/chat/completions") async def chat_completion(request: Request): data = await request.json() print_verbose(f"data passed in: {data}") @@ -186,4 +188,6 @@ async def chat_completion(request: Request): if 'stream' in data and data['stream'] == True: # use generate_responses to stream responses return StreamingResponse(data_generator(response), media_type='text/event-stream') print_verbose(f"response: {response}") - return response \ No newline at end of file + return response + +app.include_router(router) \ No newline at end of file diff --git a/litellm/tests/test_streaming.py b/litellm/tests/test_streaming.py index 1bf8eadf9da..8b2e9b9520d 100644 --- a/litellm/tests/test_streaming.py +++ b/litellm/tests/test_streaming.py @@ -709,7 +709,7 @@ def test_completion_sagemaker_stream(): except Exception as e: pytest.fail(f"Error occurred: {e}") -test_completion_sagemaker_stream() +# test_completion_sagemaker_stream() # test on openai completion call def test_openai_text_completion_call(): @@ -732,8 +732,16 @@ def test_openai_text_completion_call(): # # test on ai21 completion call def ai21_completion_call(): try: + messages=[{ + "role": "system", + "content": "You are an all-knowing oracle", + }, + { + "role": "user", + "content": "What is the meaning of the Universe?" + }] response = completion( - model="j2-ultra", messages=messages, stream=True + model="j2-ultra", messages=messages, stream=True, max_tokens=500 ) print(f"response: {response}") has_finished = False @@ -1262,3 +1270,31 @@ def test_openai_streaming_and_function_calling(): raise e # test_openai_streaming_and_function_calling() +import litellm + + +def test_success_callback_streaming(): + def success_callback(kwargs, completion_response, start_time, end_time): + print( + { + "success": True, + "input": kwargs, + "output": completion_response, + "start_time": start_time, + "end_time": end_time, + } + ) + + + litellm.success_callback = [success_callback] + + messages = [{"role": "user", "content": "hello"}] + + response = litellm.completion(model="gpt-3.5-turbo", messages=messages, stream=True) + print(response) + + + for chunk in response: + print(chunk["choices"][0]) + +test_success_callback_streaming() \ No newline at end of file diff --git a/litellm/utils.py b/litellm/utils.py index 5fb3e63b4a6..a62367b4b74 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -456,6 +456,14 @@ class Logging: end_time=end_time, print_verbose=print_verbose, ) + if callable(callback): # custom logger functions + customLogger.log_event( + kwargs=self.model_call_details, + response_obj=result, + start_time=start_time, + end_time=end_time, + print_verbose=print_verbose, + ) except Exception as e: print_verbose( @@ -2022,14 +2030,6 @@ def handle_success(args, kwargs, result, start_time, end_time): litellm_call_id=kwargs["litellm_call_id"], print_verbose=print_verbose, ) - elif callable(callback): # custom logger functions - customLogger.log_event( - kwargs=kwargs, - response_obj=result, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - ) except Exception as e: # LOGGING exception_logging(logger_fn=user_logger_fn, exception=e)