1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16 package io.netty.channel.uring;
17
18 import io.netty.channel.Channel;
19 import io.netty.channel.ChannelFuture;
20 import io.netty.channel.ChannelFutureListener;
21 import io.netty.channel.ChannelOutboundBuffer;
22 import io.netty.channel.ChannelPipeline;
23 import io.netty.channel.ChannelPromise;
24 import io.netty.channel.IoRegistration;
25 import io.netty.channel.unix.DomainSocketAddress;
26 import io.netty.channel.unix.DomainSocketChannel;
27 import io.netty.channel.unix.DomainSocketChannelConfig;
28 import io.netty.channel.unix.DomainSocketReadMode;
29 import io.netty.channel.unix.Errors;
30 import io.netty.channel.unix.FileDescriptor;
31 import io.netty.channel.unix.PeerCredentials;
32
33 import java.io.IOException;
34 import java.net.SocketAddress;
35
36
37
38
39 public final class IoUringDomainSocketChannel extends AbstractIoUringStreamChannel implements DomainSocketChannel {
40
41 private final IoUringDomainSocketChannelConfig config;
42
43 private volatile DomainSocketAddress local;
44 private volatile DomainSocketAddress remote;
45
46 public IoUringDomainSocketChannel() {
47 super(null, LinuxSocket.newSocketDomain(), false);
48 config = new IoUringDomainSocketChannelConfig(this);
49 }
50
51 IoUringDomainSocketChannel(Channel parent, FileDescriptor fd) {
52 this(parent, new LinuxSocket(fd.intValue()));
53 }
54
55 IoUringDomainSocketChannel(Channel parent, LinuxSocket fd) {
56 super(parent, fd, true);
57 local = fd.localDomainSocketAddress();
58 remote = fd.remoteDomainSocketAddress();
59 config = new IoUringDomainSocketChannelConfig(this);
60 }
61
62 @Override
63 public DomainSocketChannelConfig config() {
64 return config;
65 }
66
67 @Override
68 public DomainSocketAddress localAddress() {
69 return local;
70 }
71
72 @Override
73 public DomainSocketAddress remoteAddress() {
74 return remote;
75 }
76
77
78
79
80
81 public PeerCredentials peerCredentials() throws IOException {
82 return socket.getPeerCredentials();
83 }
84
85 @Override
86 protected Object filterOutboundMessage(Object msg) {
87 if (msg instanceof FileDescriptor) {
88 return msg;
89 }
90 return super.filterOutboundMessage(msg);
91 }
92
93 @Override
94 protected AbstractUringUnsafe newUnsafe() {
95 return new IoUringDomainSocketUnsafe();
96 }
97
98 @Override
99 protected boolean allowMultiShotPollIn() {
100
101
102 return false;
103 }
104
105 @Override
106 protected boolean socketIsEmpty(int flags) {
107 return IoUring.isUnixDomainSocketInqSupported() && super.socketIsEmpty(flags);
108 }
109
110 @Override
111 protected boolean shouldCompleteReadLoop(int flags, boolean multishot) {
112 if (IoUring.isUnixDomainSocketInqSupported()) {
113 return socketIsEmpty(flags);
114 }
115
116
117
118
119 return multishot;
120 }
121
122 private final class IoUringDomainSocketUnsafe extends IoUringStreamUnsafe {
123
124 private MsgHdrMemory writeMsgHdrMemory;
125 private MsgHdrMemory readMsgHdrMemory;
126
127 @Override
128 protected int scheduleWriteSingle(Object msg) {
129 if (msg instanceof FileDescriptor) {
130
131
132 if (writeMsgHdrMemory == null) {
133 writeMsgHdrMemory = new MsgHdrMemory();
134 }
135 IoRegistration registration = registration();
136 IoUringIoOps ioUringIoOps = prepSendFdIoOps((FileDescriptor) msg, writeMsgHdrMemory);
137 writeId = registration.submit(ioUringIoOps);
138 writeOpCode = Native.IORING_OP_SENDMSG;
139 if (writeId == 0) {
140 MsgHdrMemory memory = writeMsgHdrMemory;
141 writeMsgHdrMemory = null;
142 memory.release();
143 return 0;
144 }
145 return 1;
146 }
147 return super.scheduleWriteSingle(msg);
148 }
149
150 @Override
151 boolean writeComplete0(byte op, int res, int flags, long data, int outstanding) {
152 if (op == Native.IORING_OP_SENDMSG) {
153 writeId = 0;
154 writeOpCode = 0;
155 if (res == Native.ERRNO_ECANCELED_NEGATIVE) {
156 return true;
157 }
158 try {
159 int nativeCallResult = res >= 0 ? res : Errors.ioResult("io_uring sendmsg", res);
160 if (nativeCallResult >= 0) {
161
162 ChannelOutboundBuffer channelOutboundBuffer = unsafe().outboundBuffer();
163 if (channelOutboundBuffer != null) {
164 channelOutboundBuffer.remove();
165 }
166 }
167 } catch (Throwable throwable) {
168 handleWriteError(throwable);
169 }
170 return true;
171 }
172 return super.writeComplete0(op, res, flags, data, outstanding);
173 }
174
175 private IoUringIoOps prepSendFdIoOps(FileDescriptor fileDescriptor, MsgHdrMemory msgHdrMemory) {
176 msgHdrMemory.setScmRightsFd(fileDescriptor.intValue());
177 return IoUringIoOps.newSendmsg(
178 fd().intValue(), (byte) 0, 0, msgHdrMemory.address(), msgHdrMemory.idx());
179 }
180
181 @Override
182 protected int scheduleRead0(boolean first, boolean socketIsEmpty) {
183 DomainSocketReadMode readMode = config.getReadMode();
184 switch (readMode) {
185 case FILE_DESCRIPTORS:
186 return scheduleRecvReadFd();
187 case BYTES:
188 return super.scheduleRead0(first, socketIsEmpty);
189 default:
190 throw new Error("Unexpected read mode: " + readMode);
191 }
192 }
193
194 private int scheduleRecvReadFd() {
195
196
197 if (readMsgHdrMemory == null) {
198 readMsgHdrMemory = new MsgHdrMemory();
199 }
200 readMsgHdrMemory.prepRecvReadFd();
201 IoRegistration registration = registration();
202 IoUringIoOps ioUringIoOps = IoUringIoOps.newRecvmsg(
203 fd().intValue(), (byte) 0, 0, readMsgHdrMemory.address(), readMsgHdrMemory.idx());
204 readId = registration.submit(ioUringIoOps);
205 readOpCode = Native.IORING_OP_RECVMSG;
206 if (readId == 0) {
207 MsgHdrMemory memory = readMsgHdrMemory;
208 readMsgHdrMemory = null;
209 memory.release();
210 return 0;
211 }
212 return 1;
213 }
214
215 @Override
216 protected void readComplete0(byte op, int res, int flags, short data, int outstanding) {
217 if (op == Native.IORING_OP_RECVMSG) {
218 readId = 0;
219 if (res == Native.ERRNO_ECANCELED_NEGATIVE) {
220 return;
221 }
222 final IoUringRecvByteAllocatorHandle allocHandle = recvBufAllocHandle();
223 final ChannelPipeline pipeline = pipeline();
224 try {
225 int nativeCallResult = res >= 0 ? res : Errors.ioResult("io_uring recvmsg", res);
226 int nativeFd = readMsgHdrMemory.getScmRightsFd();
227 allocHandle.lastBytesRead(nativeFd);
228 allocHandle.incMessagesRead(1);
229 pipeline.fireChannelRead(new FileDescriptor(nativeFd));
230 } catch (Throwable throwable) {
231 handleReadException(pipeline, null, throwable, false, allocHandle);
232 } finally {
233 allocHandle.readComplete();
234 pipeline.fireChannelReadComplete();
235 }
236 return;
237 }
238 super.readComplete0(op, res, flags, data, outstanding);
239 }
240
241 @Override
242 public void connect(SocketAddress remoteAddress, SocketAddress localAddress, ChannelPromise promise) {
243
244 ChannelPromise channelPromise = newPromise().addListener(new ChannelFutureListener() {
245 @Override
246 public void operationComplete(ChannelFuture future) throws Exception {
247 if (future.isSuccess()) {
248 local = localAddress != null
249 ? (DomainSocketAddress) localAddress
250 : socket.localDomainSocketAddress();
251 remote = (DomainSocketAddress) remoteAddress;
252 promise.setSuccess();
253 } else {
254 promise.setFailure(future.cause());
255 }
256 }
257 });
258 super.connect(remoteAddress, localAddress, channelPromise);
259 }
260
261 @Override
262 public void unregistered() {
263 super.unregistered();
264 if (readMsgHdrMemory != null) {
265 readMsgHdrMemory.release();
266 readMsgHdrMemory = null;
267 }
268 if (writeMsgHdrMemory != null) {
269 writeMsgHdrMemory.release();
270 writeMsgHdrMemory = null;
271 }
272 }
273 }
274
275 @Override
276 boolean isPollInFirst() {
277 DomainSocketReadMode readMode = config.getReadMode();
278 switch (readMode) {
279 case BYTES:
280 return super.isPollInFirst();
281 case FILE_DESCRIPTORS:
282 return false;
283 default:
284 throw new Error("Unexpected read mode: " + readMode);
285 }
286 }
287 }